OpenMotionLab/MotionGPT
118
1"""OpenGL shader program wrapper.2"""3import numpy as np4import os5import re6 7import OpenGL8from OpenGL.GL import *9from OpenGL.GL import shaders as gl_shader_utils10 11 12class ShaderProgramCache(object):13 """A cache for shader programs.14 """15 16 def __init__(self, shader_dir=None):17 self._program_cache = {}18 self.shader_dir = shader_dir19 if self.shader_dir is None:20 base_dir, _ = os.path.split(os.path.realpath(__file__))21 self.shader_dir = os.path.join(base_dir, 'shaders')22 23 def get_program(self, vertex_shader, fragment_shader,24 geometry_shader=None, defines=None):25 """Get a program via a list of shader files to include in the program.26 27 Parameters28 ----------29 vertex_shader : str30 The vertex shader filename.31 fragment_shader : str32 The fragment shader filename.33 geometry_shader : str34 The geometry shader filename.35 defines : dict36 Defines and their values for the shader.37 38 Returns39 -------40 program : :class:`.ShaderProgram`41 The program.42 """43 shader_names = []44 if defines is None:45 defines = {}46 shader_filenames = [47 x for x in [vertex_shader, fragment_shader, geometry_shader]48 if x is not None49 ]50 for fn in shader_filenames:51 if fn is None:52 continue53 _, name = os.path.split(fn)54 shader_names.append(name)55 cid = OpenGL.contextdata.getContext()56 key = tuple([cid] + sorted(57 [(s,1) for s in shader_names] + [(d, defines[d]) for d in defines]58 ))59 60 if key not in self._program_cache:61 shader_filenames = [62 os.path.join(self.shader_dir, fn) for fn in shader_filenames63 ]64 if len(shader_filenames) == 2:65 shader_filenames.append(None)66 vs, fs, gs = shader_filenames67 self._program_cache[key] = ShaderProgram(68 vertex_shader=vs, fragment_shader=fs,69 geometry_shader=gs, defines=defines70 )71 return self._program_cache[key]72 73 def clear(self):74 for key in self._program_cache:75 self._program_cache[key].delete()76 self._program_cache = {}77 78 79class ShaderProgram(object):80 """A thin wrapper about OpenGL shader programs that supports easy creation,81 binding, and uniform-setting.82 83 Parameters84 ----------85 vertex_shader : str86 The vertex shader filename.87 fragment_shader : str88 The fragment shader filename.89 geometry_shader : str90 The geometry shader filename.91 defines : dict92 Defines and their values for the shader.93 """94 95 def __init__(self, vertex_shader, fragment_shader,96 geometry_shader=None, defines=None):97 98 self.vertex_shader = vertex_shader99 self.fragment_shader = fragment_shader100 self.geometry_shader = geometry_shader101 102 self.defines = defines103 if self.defines is None:104 self.defines = {}105 106 self._program_id = None107 self._vao_id = None # PYOPENGL BUG108 109 # DEBUG110 # self._unif_map = {}111 112 def _add_to_context(self):113 if self._program_id is not None:114 raise ValueError('Shader program already in context')115 shader_ids = []116 117 # Load vert shader118 shader_ids.append(gl_shader_utils.compileShader(119 self._load(self.vertex_shader), GL_VERTEX_SHADER)120 )121 # Load frag shader122 shader_ids.append(gl_shader_utils.compileShader(123 self._load(self.fragment_shader), GL_FRAGMENT_SHADER)124 )125 # Load geometry shader126 if self.geometry_shader is not None:127 shader_ids.append(gl_shader_utils.compileShader(128 self._load(self.geometry_shader), GL_GEOMETRY_SHADER)129 )130 131 # Bind empty VAO PYOPENGL BUG132 if self._vao_id is None:133 self._vao_id = glGenVertexArrays(1)134 glBindVertexArray(self._vao_id)135 136 # Compile program137 self._program_id = gl_shader_utils.compileProgram(*shader_ids)138 139 # Unbind empty VAO PYOPENGL BUG140 glBindVertexArray(0)141 142 def _in_context(self):143 return self._program_id is not None144 145 def _remove_from_context(self):146 if self._program_id is not None:147 glDeleteProgram(self._program_id)148 glDeleteVertexArrays(1, [self._vao_id])149 self._program_id = None150 self._vao_id = None151 152 def _load(self, shader_filename):153 path, _ = os.path.split(shader_filename)154 155 with open(shader_filename) as f:156 text = f.read()157 158 def ifdef(matchobj):159 if matchobj.group(1) in self.defines:160 return '#if 1'161 else:162 return '#if 0'163 164 def ifndef(matchobj):165 if matchobj.group(1) in self.defines:166 return '#if 0'167 else:168 return '#if 1'169 170 ifdef_regex = re.compile(171 '#ifdef\\s+([a-zA-Z_][a-zA-Z_0-9]*)\\s*$', re.MULTILINE172 )173 ifndef_regex = re.compile(174 '#ifndef\\s+([a-zA-Z_][a-zA-Z_0-9]*)\\s*$', re.MULTILINE175 )176 text = re.sub(ifdef_regex, ifdef, text)177 text = re.sub(ifndef_regex, ifndef, text)178 179 for define in self.defines:180 value = str(self.defines[define])181 text = text.replace(define, value)182 183 return text184 185 def _bind(self):186 """Bind this shader program to the current OpenGL context.187 """188 if self._program_id is None:189 raise ValueError('Cannot bind program that is not in context')190 # glBindVertexArray(self._vao_id)191 glUseProgram(self._program_id)192 193 def _unbind(self):194 """Unbind this shader program from the current OpenGL context.195 """196 glUseProgram(0)197 198 def delete(self):199 """Delete this shader program from the current OpenGL context.200 """201 self._remove_from_context()202 203 def set_uniform(self, name, value, unsigned=False):204 """Set a uniform value in the current shader program.205 206 Parameters207 ----------208 name : str209 Name of the uniform to set.210 value : int, float, or ndarray211 Value to set the uniform to.212 unsigned : bool213 If True, ints will be treated as unsigned values.214 """215 try:216 # DEBUG217 # self._unif_map[name] = 1, (1,)218 loc = glGetUniformLocation(self._program_id, name)219 220 if loc == -1:221 raise ValueError('Invalid shader variable: {}'.format(name))222 223 if isinstance(value, np.ndarray):224 # DEBUG225 # self._unif_map[name] = value.size, value.shape226 if value.ndim == 1:227 if (np.issubdtype(value.dtype, np.unsignedinteger) or228 unsigned):229 dtype = 'u'230 value = value.astype(np.uint32)231 elif np.issubdtype(value.dtype, np.integer):232 dtype = 'i'233 value = value.astype(np.int32)234 else:235 dtype = 'f'236 value = value.astype(np.float32)237 self._FUNC_MAP[(value.shape[0], dtype)](loc, 1, value)238 else:239 self._FUNC_MAP[(value.shape[0], value.shape[1])](240 loc, 1, GL_TRUE, value241 )242 243 # Call correct uniform function244 elif isinstance(value, float):245 glUniform1f(loc, value)246 elif isinstance(value, int):247 if unsigned:248 glUniform1ui(loc, value)249 else:250 glUniform1i(loc, value)251 elif isinstance(value, bool):252 if unsigned:253 glUniform1ui(loc, int(value))254 else:255 glUniform1i(loc, int(value))256 else:257 raise ValueError('Invalid data type')258 except Exception:259 pass260 261 _FUNC_MAP = {262 (1,'u'): glUniform1uiv,263 (2,'u'): glUniform2uiv,264 (3,'u'): glUniform3uiv,265 (4,'u'): glUniform4uiv,266 (1,'i'): glUniform1iv,267 (2,'i'): glUniform2iv,268 (3,'i'): glUniform3iv,269 (4,'i'): glUniform4iv,270 (1,'f'): glUniform1fv,271 (2,'f'): glUniform2fv,272 (3,'f'): glUniform3fv,273 (4,'f'): glUniform4fv,274 (2,2): glUniformMatrix2fv,275 (2,3): glUniformMatrix2x3fv,276 (2,4): glUniformMatrix2x4fv,277 (3,2): glUniformMatrix3x2fv,278 (3,3): glUniformMatrix3fv,279 (3,4): glUniformMatrix3x4fv,280 (4,2): glUniformMatrix4x2fv,281 (4,3): glUniformMatrix4x3fv,282 (4,4): glUniformMatrix4fv,283 }284 