@@ -105,6 +105,7 @@ def visit_Module(self, node):
105105 lines .append ('}) //end requirejs define' )
106106
107107 if self ._has_glsl :
108+ header .append ( 'var __shader_header__ = []' )
108109 header .append ( 'var __shader__ = []' )
109110
110111 lines = header + lines
@@ -200,6 +201,7 @@ def _visit_call_helper_var_glsl(self, node):
200201
201202
202203 def _visit_function (self , node ):
204+ is_main = node .name == 'main'
203205 return_type = None
204206 glsl = False
205207 glsl_wrapper_name = False
@@ -217,67 +219,85 @@ def _visit_function(self, node):
217219 return_type = decor .args [0 ].id
218220
219221 elif isinstance (decor , Attribute ) and isinstance (decor .value , Name ) and decor .value .id == 'gpu' :
220- assert decor .attr == 'vectorize'
221- gpu_vectorize = True
222+ if decor .attr == 'vectorize' :
223+ gpu_vectorize = True
224+ elif decor .attr == 'main' :
225+ is_main = True
222226
223227
224228 args = self .visit (node .args )
225229
226230 if glsl :
227- is_main = node .name == 'main'
228231 self ._has_glsl = True ## writes `__shader__ = []` in header
229232 lines = []
230233 x = []
231234 for i ,arg in enumerate (args ):
232235 if gpu_vectorize and arg not in args_typedefs :
233236 x .append ( 'float* %s' % arg )
234237 else :
235- assert arg in args_typedefs
236- x .append ( '%s %s' % (args_typedefs [arg ].replace ('POINTER' , '*' ), arg ) )
238+ if arg in args_typedefs :
239+ x .append ( '%s %s' % (args_typedefs [arg ].replace ('POINTER' , '*' ), arg ) )
240+ else :
241+ x .append ( 'float* %s' % arg )
237242
238- if is_main or gpu_vectorize :
243+ if is_main :
239244 lines .append ( '__shader__.push("void main( %s ) {");' % ', ' .join (x ) )
240245 elif return_type :
241- lines .append ( '__shader__ .push("%s %s( %s ) {");' % (return_type , node .name , ', ' .join (x )) )
246+ lines .append ( '__shader_header__ .push("%s %s( %s ) {");' % (return_type , node .name , ', ' .join (x )) )
242247 else :
243- lines .append ( '__shader__ .push("void %s( %s ) {");' % (node .name , ', ' .join (x )) )
248+ lines .append ( '__shader_header__ .push("void %s( %s ) {");' % (node .name , ', ' .join (x )) )
244249
245250 self .push ()
246- if gpu_vectorize :
251+ # `_id_` always write out an array of floats or array of vec4floats
252+ if is_main :
247253 lines .append ( '__shader__.push("vec2 _id_ = get_global_id();");' )
254+ else :
255+ lines .append ( '__shader_header__.push("vec2 _id_ = get_global_id();");' )
248256
249257 self ._glsl = True
250258 for child in node .body :
251259 if isinstance (child , Str ):
252260 continue
253261 else :
254262 for sub in self .visit (child ).splitlines ():
255- lines .append ( '__shader__.push("%s");' % (self .indent ()+ sub ) )
263+ if is_main :
264+ lines .append ( '__shader__.push("%s");' % (self .indent ()+ sub ) )
265+ else :
266+ lines .append ( '__shader_header__.push("%s");' % (self .indent ()+ sub ) )
256267 self ._glsl = False
257268 #buffer += '\n'.join(body)
258269 self .pull ()
259- lines .append ('__shader__.push("%s}");' % self .indent ())
270+ if is_main :
271+ lines .append ('__shader__.push("%s}");' % self .indent ())
272+ else :
273+ lines .append ('__shader_header__.push("%s}");' % self .indent ())
260274
261- if is_main or gpu_vectorize :
275+ lines .append (';' )
276+
277+ if is_main :
262278 if not glsl_wrapper_name :
263279 glsl_wrapper_name = node .name
264280 lines .append ('function %s( %s, __offset ) {' % (glsl_wrapper_name , ',' .join (args )) )
265- lines .append (' __offset = __offset || 0' ) ## note by default: 0 allows 0-1.0
281+ lines .append (' __offset = __offset || 0' ) ## note by default: 0 allows 0-1.0 ## TODO this needs to be set per-buffer
282+
266283 lines .append (' var __webclgl = new WebCLGL()' )
267- lines .append (' var __kernel = __webclgl.createKernel( "\\ n".join(__shader__) );' )
268- lines .append (' var __return_length = 1' )
284+ lines .append (' var header = "\\ n".join(__shader_header__)' )
285+ lines .append (' var shader = "\\ n".join(__shader__)' )
286+ lines .append (' var __kernel = __webclgl.createKernel( shader, header );' )
287+
288+ lines .append (' var __return_length = 64' ) ## minimum size is 64
269289
270290 for i ,arg in enumerate (args ):
271291 lines .append (' if (%s instanceof Array) {' % arg )
272- lines .append (' __return_length = %s.length' % arg )
292+ lines .append (' __return_length = %s.length==2 ? %s : %s.length ' % ( arg , arg , arg ) )
273293 lines .append (' var %s_buffer = __webclgl.createBuffer(%s.length, "FLOAT", __offset)' % (arg ,arg ))
274294 lines .append (' __webclgl.enqueueWriteBuffer(%s_buffer, %s)' % (arg , arg ))
275295 lines .append (' __kernel.setKernelArg(%s, %s_buffer)' % (i , arg ))
276296 lines .append (' } else { __kernel.setKernelArg(%s, %s) }' % (i , arg ))
277297
278298 lines .append (' var return_buffer = __webclgl.createBuffer(__return_length, "FLOAT", __offset)' )
279-
280299 lines .append (' __kernel.compile()' )
300+
281301 lines .append (' __webclgl.enqueueNDRangeKernel(__kernel, return_buffer)' )
282302 lines .append (' return __webclgl.enqueueReadBuffer_Float( return_buffer )' )
283303 lines .append ('} // end of wrapper' )
0 commit comments