MayaFlux 0.5.0
Digital-First Multimedia Processing Framework
Loading...
Searching...
No Matches
AsmGenerator.cpp
Go to the documentation of this file.
1#include "ShaderFoundry.hpp"
2
3/**
4 * @file AsmGenerator.cpp
5 * @brief SPIR-V assembly and GLSL kernel emitters for ShaderSpec.
6 *
7 * The assembly path (emit_spirv_asm) covers all binding modalities and
8 * kernel templates:
9 *
10 * SSBO (SCALAR_F32, VEC2_F32, VEC3_F32, VEC4_F32) — all KernelOps via
11 * emit_elementwise_body. ArrayStride, element type, and pointer type derived
12 * from GpuDataFormat. Scalar PC operands splatted to vector width via
13 * OpCompositeConstruct before arithmetic.
14 *
15 * IMAGE_2D storage images — emit_image_body handles any number of Input/InOut
16 * images and one Output image with all single- and two-operand KernelOps
17 * applied per channel.
18 *
19 * TEXTURE_2D sampled images — emit_image_body loads each binding via
20 * OpTypeSampledImage / OpImageSampleExplicitLod at normalised UV derived
21 * from GlobalInvocationID divided by width/height PC fields (first two PC
22 * fields by convention when TEXTURE_2D bindings are present).
23 *
24 * KernelTemplate::Reduction — emit_reduction_body emits a structured
25 * workgroup tree reduction with OpLoopMerge, OpPhi for the stride variable,
26 * OpULessThan active-thread guard, and two OpControlBarrier synchronisation
27 * points per iteration. Supports KernelOp::Sum and KernelOp::Max.
28 *
29 * GLSL kernel path (emit_glsl_kernel) handles all three binding modalities
30 * when spec.kernel is set via MF_KERNEL.
31 */
32
34
35namespace {
36
37 /**
38 * Maps GpuDataFormat to the GLSL type name used in SSBO array declarations.
39 */
40 std::string_view glsl_type(Kakshya::GpuDataFormat fmt) noexcept
41 {
42 switch (fmt) {
44 return "float";
46 return "vec2";
48 return "vec3";
50 return "vec4";
52 return "int";
56 return "uint";
57 default:
58 return "float";
59 }
60 }
61
62 /**
63 * Returns the SPIR-V type ID string for the element type of an SSBO binding.
64 */
65 std::string_view ssbo_elem_spirv_type(Kakshya::GpuDataFormat fmt) noexcept
66 {
67 switch (fmt) {
69 return "%v2f32";
71 return "%v3f32";
73 return "%v4f32";
77 return "%u32";
79 return "%i32";
80 default:
81 return "%f32";
82 }
83 }
84
85 /**
86 * Returns the component count for a GpuDataFormat SSBO element.
87 * Scalar formats return 1.
88 */
89 uint32_t ssbo_elem_components(Kakshya::GpuDataFormat fmt) noexcept
90 {
91 switch (fmt) {
93 return 2;
95 return 3;
97 return 4;
98 default:
99 return 1;
100 }
101 }
102
103 /**
104 * Emit a fixed header common to all generated compute kernels.
105 * Assigns IDs for void, voidfn, u32, f32, v3u32, glsl extension import,
106 * and the GlobalInvocationId builtin.
107 */
108 std::string emit_header(const ShaderSpec& spec)
109 {
110 const auto& ws = spec.workgroup_size;
111
112 bool has_storage_image = false;
113 bool has_input_image = false;
114 for (const auto& b : spec.bindings) {
115 if (b.modality == Kakshya::DataModality::IMAGE_2D) {
116 has_storage_image = true;
117 if (b.direction != BindingDirection::Output)
118 has_input_image = true;
119 }
120 }
121
122 std::string iface = "%gid_var";
123 for (const auto& b : spec.bindings) {
124 if (b.modality == Kakshya::DataModality::IMAGE_2D)
125 continue;
126 if (b.modality == Kakshya::DataModality::TEXTURE_2D) {
127 iface += " %tex_" + b.name;
128 continue;
129 }
130 iface += " %buf_" + b.name;
131 }
132 if (has_storage_image) {
133 for (const auto& b : spec.bindings) {
134 if (b.modality == Kakshya::DataModality::IMAGE_2D)
135 iface += " %img_" + b.name;
136 }
137 }
138 if (!spec.pc_fields.empty())
139 iface += " %pc";
140
142 iface += " %lid_var";
143
144 std::string o;
145 o += "OpCapability Shader\n";
146
147 if (has_storage_image)
148 o += "OpCapability StorageImageWriteWithoutFormat\n";
149 if (has_input_image)
150 o += "OpCapability StorageImageReadWithoutFormat\n";
151
152 o += "%glsl = OpExtInstImport \"GLSL.std.450\"\n";
153 o += "OpMemoryModel Logical GLSL450\n";
154 o += "OpEntryPoint GLCompute %main \"main\" " + iface + "\n";
155 o += "OpExecutionMode %main LocalSize "
156 + std::to_string(ws[0]) + " "
157 + std::to_string(ws[1]) + " "
158 + std::to_string(ws[2]) + "\n\n";
159 return o;
160 }
161
162 /**
163 * Emit decorations for all SSBO bindings and push constant member offsets.
164 */
165 std::string emit_decorations(const ShaderSpec& spec)
166 {
167 std::string o;
168 o += "OpDecorate %gid_var BuiltIn GlobalInvocationId\n";
169
170 for (const auto& b : spec.bindings) {
172 || b.modality == Kakshya::DataModality::IMAGE_2D)
173 continue;
174
175 const auto stride = static_cast<uint32_t>(Kakshya::gpu_data_format_bytes(b.format));
176 o += "OpDecorate %rta_" + b.name + " ArrayStride "
177 + std::to_string(stride) + "\n";
178 o += "OpMemberDecorate %blk_" + b.name + " 0 Offset 0\n";
179 o += "OpDecorate %blk_" + b.name + " Block\n";
180 o += "OpDecorate %buf_" + b.name + " DescriptorSet " + std::to_string(b.set) + "\n";
181 o += "OpDecorate %buf_" + b.name + " Binding "
182 + std::to_string(b.binding_index) + "\n";
183 }
184
185 for (const auto& b : spec.bindings) {
186 if (b.modality != Kakshya::DataModality::IMAGE_2D)
187 continue;
188
189 const std::string var = "%img_" + b.name;
190 o += "OpDecorate " + var + " DescriptorSet " + std::to_string(b.set) + "\n";
191 o += "OpDecorate " + var + " Binding "
192 + std::to_string(b.binding_index) + "\n";
193
194 if (b.direction == BindingDirection::Input) {
195 o += "OpDecorate " + var + " NonWritable\n";
196 } else if (b.direction == BindingDirection::Output) {
197 o += "OpDecorate " + var + " NonReadable\n";
198 }
199 }
200
201 for (const auto& b : spec.bindings) {
202 if (b.modality != Kakshya::DataModality::TEXTURE_2D)
203 continue;
204 o += "OpDecorate %tex_" + b.name + " DescriptorSet " + std::to_string(b.set) + "\n";
205 o += "OpDecorate %tex_" + b.name + " Binding "
206 + std::to_string(b.binding_index) + "\n";
207 }
208
210 o += "OpDecorate %lid_var BuiltIn LocalInvocationId\n";
211
212 if (!spec.pc_fields.empty()) {
213 o += "OpDecorate %pc_blk Block\n";
214 uint32_t off = 0;
215 for (size_t i = 0; i < spec.pc_fields.size(); ++i) {
216 o += "OpMemberDecorate %pc_blk " + std::to_string(i)
217 + " Offset " + std::to_string(off) + "\n";
218 off += static_cast<uint32_t>(
220 }
221 }
222 o += "\n";
223 return o;
224 }
225
226 /**
227 * Emit type declarations for all used types, including void, voidfn, u32, f32,
228 * v3u32, and the GlobalInvocationId builtin.
229 */
230 std::string emit_types(const ShaderSpec& spec)
231 {
232 std::string o;
233 o += "%void = OpTypeVoid\n";
234 o += "%voidfn = OpTypeFunction %void\n";
235 o += "%u32 = OpTypeInt 32 0\n";
236 o += "%f32 = OpTypeFloat 32\n";
237 o += "%v3u32 = OpTypeVector %u32 3\n";
238 o += "%ptr_in_v3u32 = OpTypePointer Input %v3u32\n";
239 o += "%gid_var = OpVariable %ptr_in_v3u32 Input\n\n";
240 o += "%bool = OpTypeBool\n";
241
242 bool need_v2f32 = false;
243 bool need_v3f32 = false;
244 bool need_v4f32 = false;
245
246 for (const auto& b : spec.bindings) {
248 || b.modality == Kakshya::DataModality::IMAGE_2D)
249 continue;
250 switch (b.format) {
252 need_v2f32 = true;
253 break;
255 need_v3f32 = true;
256 break;
258 need_v4f32 = true;
259 break;
260 default:
261 break;
262 }
263 }
264
265 if (need_v2f32)
266 o += "%v2f32 = OpTypeVector %f32 2\n";
267 if (need_v3f32)
268 o += "%v3f32 = OpTypeVector %f32 3\n";
269
270 bool need_i32 = false;
271 for (const auto& b : spec.bindings) {
273 || b.modality == Kakshya::DataModality::IMAGE_2D)
274 continue;
275 if (b.format == Kakshya::GpuDataFormat::INT32)
276 need_i32 = true;
277 }
278
279 bool has_image_2d = false;
280 for (const auto& b : spec.bindings) {
281 if (b.modality == Kakshya::DataModality::IMAGE_2D) {
282 has_image_2d = true;
283 break;
284 }
285 }
286 if (need_v4f32 && !has_image_2d)
287 o += "%v4f32 = OpTypeVector %f32 4\n";
288 if (need_v2f32 || need_v3f32 || (need_v4f32 && !has_image_2d))
289 o += "\n";
290 if (need_i32 && !has_image_2d)
291 o += "%i32 = OpTypeInt 32 1\n";
292
293 for (const auto& b : spec.bindings) {
295 || b.modality == Kakshya::DataModality::IMAGE_2D)
296 continue;
297
298 const std::string_view etype = ssbo_elem_spirv_type(b.format);
299 o += "%rta_" + b.name + " = OpTypeRuntimeArray " + std::string(etype) + "\n";
300 o += "%blk_" + b.name + " = OpTypeStruct %rta_" + b.name + "\n";
301 o += "%pblk_" + b.name + " = OpTypePointer StorageBuffer %blk_" + b.name + "\n";
302 o += "%buf_" + b.name + " = OpVariable %pblk_" + b.name + " StorageBuffer\n";
303 o += "%pelem_" + b.name + " = OpTypePointer StorageBuffer "
304 + std::string(etype) + "\n";
305 }
306 o += "\n";
307
308 if (has_image_2d) {
309 o += "%i32 = OpTypeInt 32 1\n";
310 o += "%v2i32 = OpTypeVector %i32 2\n";
311 o += "%v4f32 = OpTypeVector %f32 4\n";
312 o += "%img_czero = OpConstant %f32 0.0\n";
313 o += "%img_cone = OpConstant %f32 1.0\n";
314 o += "%ci0 = OpConstant %i32 0\n";
315 o += "%ci1 = OpConstant %i32 1\n";
316 o += "%img2d_t = OpTypeImage %f32 2D 0 0 0 2 Unknown\n";
317 o += "%ptr_img2d = OpTypePointer UniformConstant %img2d_t\n";
318 for (const auto& b : spec.bindings) {
319 if (b.modality != Kakshya::DataModality::IMAGE_2D)
320 continue;
321 o += "%img_" + b.name + " = OpVariable %ptr_img2d UniformConstant\n";
322 }
323 o += "\n";
324 }
325
326 bool has_texture_2d = false;
327 for (const auto& b : spec.bindings) {
328 if (b.modality == Kakshya::DataModality::TEXTURE_2D) {
329 has_texture_2d = true;
330 break;
331 }
332 }
333
334 if (has_texture_2d) {
335 if (!need_v2f32 && !has_image_2d)
336 o += "%v2f32 = OpTypeVector %f32 2\n";
337 o += "%sampler_t = OpTypeSampler\n";
338 o += "%img2d_s_t = OpTypeImage %f32 2D 0 0 0 1 Unknown\n";
339 o += "%simgc_t = OpTypeSampledImage %img2d_s_t\n";
340 o += "%ptr_simg = OpTypePointer UniformConstant %simgc_t\n";
341 for (const auto& b : spec.bindings) {
342 if (b.modality != Kakshya::DataModality::TEXTURE_2D)
343 continue;
344 o += "%tex_" + b.name + " = OpVariable %ptr_simg UniformConstant\n";
345 }
346 o += "%lod_zero = OpConstant %f32 0.0\n";
347 o += "\n";
348 }
349
350 if (spec.tmpl == KernelTemplate::Reduction) {
351 const uint32_t local = spec.workgroup_size[0];
352 const std::string ls = std::to_string(local);
353 const bool is_max_index = (spec.op == KernelOp::MaxIndex);
354
355 o += "%bool = OpTypeBool\n";
356 o += "%lid_var = OpVariable %ptr_in_v3u32 Input\n";
357 o += "%arr_sh_t = OpTypeArray %f32 %c" + ls + "u\n";
358 o += "%psh_arr = OpTypePointer Workgroup %arr_sh_t\n";
359 o += "%shared_a = OpVariable %psh_arr Workgroup\n";
360 o += "%psh_f32 = OpTypePointer Workgroup %f32\n";
361
362 if (is_max_index) {
363 o += "%arr_shu_t = OpTypeArray %u32 %c" + ls + "u\n";
364 o += "%pshu_arr = OpTypePointer Workgroup %arr_shu_t\n";
365 o += "%shared_idx = OpVariable %pshu_arr Workgroup\n";
366 o += "%psh_u32 = OpTypePointer Workgroup %u32\n";
367 }
368
369 o += "%c" + ls + "u = OpConstant %u32 " + ls + "\n";
370 o += "%c2u = OpConstant %u32 2\n";
371 o += "%c264u = OpConstant %u32 264\n";
372 o += "\n";
373 } else if (spec.tmpl == KernelTemplate::Scan) {
374 const uint32_t local = spec.workgroup_size[0];
375 const std::string ls = std::to_string(local);
376
377 o += "%bool = OpTypeBool\n";
378 o += "%lid_var = OpVariable %ptr_in_v3u32 Input\n";
379 o += "%arr_sh_t = OpTypeArray %f32 %c" + ls + "u\n";
380 o += "%psh_arr = OpTypePointer Workgroup %arr_sh_t\n";
381 o += "%shared_a = OpVariable %psh_arr Workgroup\n";
382 o += "%shared_b = OpVariable %psh_arr Workgroup\n";
383 o += "%psh_f32 = OpTypePointer Workgroup %f32\n";
384 o += "%c" + ls + "u = OpConstant %u32 " + ls + "\n";
385 o += "%c2u = OpConstant %u32 2\n";
386 o += "%c264u = OpConstant %u32 264\n";
387 o += "\n";
388 }
389
390 if (!spec.pc_fields.empty()) {
391 bool pc_has_uint = false;
392 bool pc_has_int = false;
393 o += "%pc_blk = OpTypeStruct";
394 for (const auto& f : spec.pc_fields) {
395 o += " " + std::string(ssbo_elem_spirv_type(f.format));
396 if (f.format == Kakshya::GpuDataFormat::UINT32) {
397 pc_has_uint = true;
398 } else if (f.format == Kakshya::GpuDataFormat::INT32) {
399 pc_has_int = true;
400 }
401 }
402 o += "\n";
403 o += "%ppc = OpTypePointer PushConstant %pc_blk\n";
404 o += "%pc = OpVariable %ppc PushConstant\n";
405 o += "%ppc_f32 = OpTypePointer PushConstant %f32\n";
406 if (pc_has_uint)
407 o += "%ppc_u32 = OpTypePointer PushConstant %u32\n";
408 if (pc_has_int)
409 o += "%ppc_i32 = OpTypePointer PushConstant %i32\n";
410 o += "\n";
411 }
412
413 o += "%c0u = OpConstant %u32 0\n";
414 o += "%c1u = OpConstant %u32 1\n";
415
416 for (size_t i = 2; i < spec.pc_fields.size(); ++i) {
417 o += "%c" + std::to_string(i) + "u = OpConstant %u32 "
418 + std::to_string(i) + "\n";
419 }
420
421 o += "\n";
422 return o;
423 }
424
425 /**
426 * Emit the entry point body for Elementwise/Stencil templates.
427 * Loads the index, loads PC fields, loads SSBO elements, applies op, stores.
428 */
429 std::string emit_elementwise_body(const ShaderSpec& spec)
430 {
431 std::string o;
432 o += "%main = OpFunction %void None %voidfn\n";
433 o += "%entry = OpLabel\n";
434 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
435 o += "%i = OpCompositeExtract %u32 %gid3 0\n\n";
436
437 for (size_t fi = 0; fi < spec.pc_fields.size(); ++fi) {
438 const auto& f = spec.pc_fields[fi];
439 const std::string_view pptr = (f.format == Kakshya::GpuDataFormat::UINT32)
440 ? "%ppc_u32"
441 : (f.format == Kakshya::GpuDataFormat::INT32 ? "%ppc_i32" : "%ppc_f32");
442 const std::string_view ld_type = ssbo_elem_spirv_type(f.format);
443 o += "%ppc_" + f.name + " = OpAccessChain " + std::string(pptr)
444 + " %pc %c" + std::to_string(fi) + "u\n";
445 o += "%pc_" + f.name + " = OpLoad " + std::string(ld_type)
446 + " %ppc_" + f.name + "\n";
447 }
448 o += "\n";
449
450 std::vector<const BindingSlot*> ssbos;
451 for (const auto& b : spec.bindings) {
453 || b.modality == Kakshya::DataModality::IMAGE_2D)
454 continue;
455 ssbos.push_back(&b);
456 }
457
459 for (const auto* b : ssbos) {
460 if (b->direction != BindingDirection::Output) {
461 primary_fmt = b->format;
462 break;
463 }
464 }
465 const std::string_view etype = ssbo_elem_spirv_type(primary_fmt);
466 const uint32_t ncomp = ssbo_elem_components(primary_fmt);
467 const bool is_vector = (ncomp > 1);
468
469 for (const auto* b : ssbos) {
470 o += "%gep_" + b->name + " = OpAccessChain %pelem_" + b->name
471 + " %buf_" + b->name + " %c0u %i\n";
472 if (b->direction != BindingDirection::Output) {
473 o += "%val_" + b->name + " = OpLoad " + std::string(etype)
474 + " %gep_" + b->name + "\n";
475 }
476 }
477 o += "\n";
478
479 auto pc_operand = [&](const std::string& field_name) -> std::string {
480 if (!is_vector)
481 return "%pc_" + field_name;
482 const std::string splat = "%spc_" + field_name;
483 std::string construct = splat + " = OpCompositeConstruct "
484 + std::string(etype);
485 const std::string scalar = " %pc_" + field_name;
486 for (uint32_t c = 0; c < ncomp; ++c)
487 construct += scalar;
488 o += construct + "\n";
489 return splat;
490 };
491
492 const std::string v0 = ssbos.empty() ? "" : ("%val_" + ssbos[0]->name);
493 const std::string v1 = ssbos.size() > 1 ? ("%val_" + ssbos[1]->name) : "";
494 const std::string p0 = spec.pc_fields.empty() ? ""
495 : pc_operand(spec.pc_fields[0].name);
496 const std::string p1 = spec.pc_fields.size() > 1
497 ? pc_operand(spec.pc_fields[1].name)
498 : "";
499
500 std::string result;
501 const std::string et = std::string(etype);
502 switch (spec.op) {
503 case KernelOp::Scale:
504 if (is_vector) {
505 o += "%res = OpVectorTimesScalar " + et + " " + v0
506 + " %pc_" + spec.pc_fields[0].name + "\n";
507 } else {
508 o += "%res = OpFMul " + et + " " + v0 + " " + p0 + "\n";
509 }
510
511 result = "%res";
512 break;
514 if (is_vector) {
515 o += "%scaled = OpVectorTimesScalar " + et + " " + v0
516 + " %pc_" + spec.pc_fields[0].name + "\n";
517 o += "%res = OpFAdd " + et + " %scaled " + p1 + "\n";
518 } else {
519 o += "%mul = OpFMul " + et + " " + v0 + " " + p0 + "\n";
520 o += "%res = OpFAdd " + et + " %mul " + p1 + "\n";
521 }
522
523 result = "%res";
524 break;
525 case KernelOp::Offset:
526 o += "%res = OpFAdd " + et + " " + v0 + " " + p0 + "\n";
527 result = "%res";
528 break;
529 case KernelOp::Clip:
530 o += "%res = OpExtInst " + et + " %glsl FClamp " + v0
531 + " " + p0 + " " + p1 + "\n";
532 result = "%res";
533 break;
534 case KernelOp::Abs:
535 o += "%res = OpExtInst " + et + " %glsl FAbs " + v0 + "\n";
536 result = "%res";
537 break;
538 case KernelOp::Negate:
539 o += "%res = OpFNegate " + et + " " + v0 + "\n";
540 result = "%res";
541 break;
542 case KernelOp::Add:
543 o += "%res = OpFAdd " + et + " " + v0 + " " + v1 + "\n";
544 result = "%res";
545 break;
547 o += "%res = OpFMul " + et + " " + v0 + " " + v1 + "\n";
548 result = "%res";
549 break;
550 case KernelOp::Mix:
551 o += "%dlt = OpFSub %f32 " + v1 + " " + v0 + "\n";
552 o += "%scl = OpFMul %f32 %dlt " + p0 + "\n";
553 o += "%res = OpFAdd %f32 " + v0 + " %scl\n";
554 result = "%res";
555 break;
556 case KernelOp::Sub:
557 o += "%res = OpFSub " + et + " " + v0 + " " + v1 + "\n";
558 result = "%res";
559 break;
560 case KernelOp::Fma:
561 o += "%res = OpExtInst %f32 %glsl Fma " + v0 + " " + p0 + " " + p1 + "\n";
562 result = "%res";
563 break;
564 case KernelOp::Floor:
565 o += "%res = OpExtInst %f32 %glsl Floor " + v0 + "\n";
566 result = "%res";
567 break;
568 case KernelOp::Ceil:
569 o += "%res = OpExtInst %f32 %glsl Ceil " + v0 + "\n";
570 result = "%res";
571 break;
572 case KernelOp::Round:
573 o += "%res = OpExtInst %f32 %glsl Round " + v0 + "\n";
574 result = "%res";
575 break;
576 case KernelOp::Trunc:
577 o += "%res = OpExtInst %f32 %glsl Trunc " + v0 + "\n";
578 result = "%res";
579 break;
580 case KernelOp::Fract:
581 o += "%res = OpExtInst %f32 %glsl Fract " + v0 + "\n";
582 result = "%res";
583 break;
584 case KernelOp::Sqrt:
585 o += "%res = OpExtInst %f32 %glsl Sqrt " + v0 + "\n";
586 result = "%res";
587 break;
589 o += "%res = OpExtInst %f32 %glsl InverseSqrt " + v0 + "\n";
590 result = "%res";
591 break;
592 case KernelOp::Sin:
593 o += "%res = OpExtInst %f32 %glsl Sin " + v0 + "\n";
594 result = "%res";
595 break;
596 case KernelOp::Cos:
597 o += "%res = OpExtInst %f32 %glsl Cos " + v0 + "\n";
598 result = "%res";
599 break;
600 case KernelOp::Tan:
601 o += "%res = OpExtInst %f32 %glsl Tan " + v0 + "\n";
602 result = "%res";
603 break;
604 case KernelOp::Asin:
605 o += "%res = OpExtInst %f32 %glsl Asin " + v0 + "\n";
606 result = "%res";
607 break;
608 case KernelOp::Acos:
609 o += "%res = OpExtInst %f32 %glsl Acos " + v0 + "\n";
610 result = "%res";
611 break;
612 case KernelOp::Atan:
613 o += "%res = OpExtInst %f32 %glsl Atan " + v0 + "\n";
614 result = "%res";
615 break;
616 case KernelOp::Sinh:
617 o += "%res = OpExtInst %f32 %glsl Sinh " + v0 + "\n";
618 result = "%res";
619 break;
620 case KernelOp::Cosh:
621 o += "%res = OpExtInst %f32 %glsl Cosh " + v0 + "\n";
622 result = "%res";
623 break;
624 case KernelOp::Tanh:
625 o += "%res = OpExtInst %f32 %glsl Tanh " + v0 + "\n";
626 result = "%res";
627 break;
628 case KernelOp::Exp:
629 o += "%res = OpExtInst %f32 %glsl Exp " + v0 + "\n";
630 result = "%res";
631 break;
632 case KernelOp::Exp2:
633 o += "%res = OpExtInst %f32 %glsl Exp2 " + v0 + "\n";
634 result = "%res";
635 break;
636 case KernelOp::Log:
637 o += "%res = OpExtInst %f32 %glsl Log " + v0 + "\n";
638 result = "%res";
639 break;
640 case KernelOp::Log2:
641 o += "%res = OpExtInst %f32 %glsl Log2 " + v0 + "\n";
642 result = "%res";
643 break;
644 case KernelOp::Pow:
645 o += "%res = OpExtInst %f32 %glsl Pow " + v0 + " " + v1 + "\n";
646 result = "%res";
647 break;
648 case KernelOp::Atan2:
649 o += "%res = OpExtInst %f32 %glsl Atan2 " + v0 + " " + v1 + "\n";
650 result = "%res";
651 break;
652 case KernelOp::Min:
653 o += "%res = OpExtInst %f32 %glsl FMin " + v0 + " " + v1 + "\n";
654 result = "%res";
655 break;
656 case KernelOp::MaxTwo:
657 o += "%res = OpExtInst %f32 %glsl FMax " + v0 + " " + v1 + "\n";
658 result = "%res";
659 break;
660 case KernelOp::Step:
661 o += "%res = OpExtInst %f32 %glsl Step " + v0 + " " + v1 + "\n";
662 result = "%res";
663 break;
665 o += "%res = OpExtInst %f32 %glsl SmoothStep " + v0 + " " + v1 + " " + p0 + "\n";
666 result = "%res";
667 break;
669 o += "%i_f32 = OpConvertUToF %f32 %i\n";
670 o += "%res = OpFMul " + et + " " + v0 + " %i_f32\n";
671 result = "%res";
672 break;
673 }
674 default:
675 o += "%res = OpCopyObject " + et + " " + v0 + "\n";
676 result = "%res";
677 break;
678 }
679 o += "\n";
680
681 for (const auto* b : ssbos) {
682 if (b->direction == BindingDirection::Input)
683 continue;
684 o += "OpStore %gep_" + b->name + " " + result + "\n";
685 }
686
687 o += "OpReturn\n";
688 o += "OpFunctionEnd\n";
689 return o;
690 }
691
692 std::string emit_reduction_body(const ShaderSpec& spec)
693 {
694 const uint32_t local = spec.workgroup_size[0];
695 const std::string ls = std::to_string(local);
696 const bool is_max = (spec.op == KernelOp::Max);
697 const auto& b0 = spec.bindings.front();
698
699 std::string o;
700 o += "%main = OpFunction %void None %voidfn\n";
701 o += "%entry = OpLabel\n";
702
703 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
704 o += "%i = OpCompositeExtract %u32 %gid3 0\n";
705 o += "%lid3 = OpLoad %v3u32 %lid_var\n";
706 o += "%lid = OpCompositeExtract %u32 %lid3 0\n\n";
707
708 o += "%gep_in = OpAccessChain %pelem_" + b0.name
709 + " %buf_" + b0.name + " %c0u %i\n";
710 o += "%elem = OpLoad %f32 %gep_in\n";
711 o += "%pgsh = OpAccessChain %psh_f32 %shared_a %lid\n";
712 o += "OpStore %pgsh %elem\n";
713 o += "OpControlBarrier %c2u %c2u %c264u\n\n";
714
715 o += "%s_init = OpShiftRightLogical %u32 %c" + ls + "u %c1u\n";
716 o += "OpBranch %loop_hdr\n\n";
717
718 o += "%loop_hdr = OpLabel\n";
719 o += "%stride = OpPhi %u32 %s_init %entry %s_next %loop_cont\n";
720 o += "OpLoopMerge %loop_merge %loop_cont None\n";
721 o += "OpBranch %loop_body\n\n";
722
723 o += "%loop_body = OpLabel\n";
724 o += "%active = OpULessThan %bool %lid %stride\n";
725 o += "OpSelectionMerge %sel_merge None\n";
726 o += "OpBranchConditional %active %do_op %sel_merge\n\n";
727
728 o += "%do_op = OpLabel\n";
729 o += "%pgsh_a = OpAccessChain %psh_f32 %shared_a %lid\n";
730 o += "%a = OpLoad %f32 %pgsh_a\n";
731 o += "%lid_b = OpIAdd %u32 %lid %stride\n";
732 o += "%pgsh_b = OpAccessChain %psh_f32 %shared_a %lid_b\n";
733 o += "%b = OpLoad %f32 %pgsh_b\n";
734
735 if (is_max) {
736 o += "%combined = OpExtInst %f32 %glsl FMax %a %b\n";
737 } else {
738 o += "%combined = OpFAdd %f32 %a %b\n";
739 }
740
741 o += "OpStore %pgsh_a %combined\n";
742 o += "OpBranch %sel_merge\n\n";
743
744 o += "%sel_merge = OpLabel\n";
745 o += "OpBranch %loop_cont\n\n";
746
747 o += "%loop_cont = OpLabel\n";
748 o += "OpControlBarrier %c2u %c2u %c264u\n";
749 o += "%s_next = OpShiftRightLogical %u32 %stride %c1u\n";
750 o += "%done = OpIEqual %bool %s_next %c0u\n";
751 o += "OpBranchConditional %done %loop_merge %loop_hdr\n\n";
752
753 o += "%loop_merge = OpLabel\n";
754 o += "%is_zero = OpIEqual %bool %lid %c0u\n";
755 o += "OpSelectionMerge %write_merge None\n";
756 o += "OpBranchConditional %is_zero %do_write %write_merge\n\n";
757
758 o += "%do_write = OpLabel\n";
759 o += "%result = OpLoad %f32 %pgsh\n";
760 o += "%gep_out = OpAccessChain %pelem_" + b0.name
761 + " %buf_" + b0.name + " %c0u %c0u\n";
762 o += "OpStore %gep_out %result\n";
763 o += "OpBranch %write_merge\n\n";
764
765 o += "%write_merge = OpLabel\n";
766 o += "OpReturn\n";
767 o += "OpFunctionEnd\n";
768 return o;
769 }
770
771 /**
772 * @brief Emit the entry point body for the Scan template (inclusive
773 * prefix sum over one InOut SSBO).
774 *
775 * Double-buffered Hillis-Steele scan in workgroup shared memory:
776 * log2(local_size) fixed passes, unrolled at generation time, each pass
777 * reading exclusively from one shared array and writing exclusively to
778 * the other, so no lane can read a value another lane has already
779 * overwritten in the same pass. Strides are derived via successive
780 * OpShiftLeftLogical on %c1u rather than emitting a fresh OpConstant
781 * per pass, since a literal-valued constant can collide with an
782 * existing constant of the same value already declared elsewhere in
783 * the module.
784 *
785 * @param spec ShaderSpec with tmpl == KernelTemplate::Scan and exactly
786 * one InOut FLOAT32 SSBO binding.
787 * @return SPIR-V function body text for %main.
788 */
789 std::string emit_scan_body(const ShaderSpec& spec)
790 {
791 const uint32_t local = spec.workgroup_size[0];
792 const auto log2_local = static_cast<uint32_t>(std::log2(static_cast<double>(local)));
793 const auto& b0 = spec.bindings.front();
794
795 std::string o;
796 o += "%main = OpFunction %void None %voidfn\n";
797 o += "%entry = OpLabel\n";
798
799 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
800 o += "%i = OpCompositeExtract %u32 %gid3 0\n";
801 o += "%lid3 = OpLoad %v3u32 %lid_var\n";
802 o += "%lid = OpCompositeExtract %u32 %lid3 0\n\n";
803
804 o += "%gep_in = OpAccessChain %pelem_" + b0.name
805 + " %buf_" + b0.name + " %c0u %i\n";
806 o += "%elem = OpLoad %f32 %gep_in\n";
807 o += "%pgsh0 = OpAccessChain %psh_f32 %shared_a %lid\n";
808 o += "OpStore %pgsh0 %elem\n";
809 o += "OpControlBarrier %c2u %c2u %c264u\n\n";
810
811 std::string read_buf = "%shared_a";
812 std::string write_buf = "%shared_b";
813 std::string stride_val = "%c1u";
814
815 for (uint32_t pass = 0; pass < log2_local; ++pass) {
816 const std::string ps = std::to_string(pass);
817
818 o += "%has_left_" + ps + " = OpUGreaterThanEqual %bool %lid " + stride_val + "\n";
819 o += "OpSelectionMerge %scan_merge_" + ps + " None\n";
820 o += "OpBranchConditional %has_left_" + ps + " %scan_add_" + ps + " %scan_copy_" + ps + "\n\n";
821
822 o += "%scan_add_" + ps + " = OpLabel\n";
823 o += "%self_ptr_" + ps + " = OpAccessChain %psh_f32 " + read_buf + " %lid\n";
824 o += "%self_val_" + ps + " = OpLoad %f32 %self_ptr_" + ps + "\n";
825 o += "%left_idx_" + ps + " = OpISub %u32 %lid " + stride_val + "\n";
826 o += "%left_ptr_" + ps + " = OpAccessChain %psh_f32 " + read_buf + " %left_idx_" + ps + "\n";
827 o += "%left_val_" + ps + " = OpLoad %f32 %left_ptr_" + ps + "\n";
828 o += "%sum_" + ps + " = OpFAdd %f32 %self_val_" + ps + " %left_val_" + ps + "\n";
829 o += "%wptr_add_" + ps + " = OpAccessChain %psh_f32 " + write_buf + " %lid\n";
830 o += "OpStore %wptr_add_" + ps + " %sum_" + ps + "\n";
831 o += "OpBranch %scan_merge_" + ps + "\n\n";
832
833 o += "%scan_copy_" + ps + " = OpLabel\n";
834 o += "%pass_ptr_" + ps + " = OpAccessChain %psh_f32 " + read_buf + " %lid\n";
835 o += "%pass_val_" + ps + " = OpLoad %f32 %pass_ptr_" + ps + "\n";
836 o += "%wptr_copy_" + ps + " = OpAccessChain %psh_f32 " + write_buf + " %lid\n";
837 o += "OpStore %wptr_copy_" + ps + " %pass_val_" + ps + "\n";
838 o += "OpBranch %scan_merge_" + ps + "\n\n";
839
840 o += "%scan_merge_" + ps + " = OpLabel\n";
841 o += "OpControlBarrier %c2u %c2u %c264u\n\n";
842
843 std::swap(read_buf, write_buf);
844
845 if (pass + 1 < log2_local) {
846 const std::string next_stride = "%stride_next_" + ps;
847 o += next_stride + " = OpShiftLeftLogical %u32 " + stride_val + " %c1u\n";
848 stride_val = next_stride;
849 }
850 }
851
852 o += "%final_ptr = OpAccessChain %psh_f32 " + read_buf + " %lid\n";
853 o += "%final_val = OpLoad %f32 %final_ptr\n";
854 o += "%gep_out = OpAccessChain %pelem_" + b0.name
855 + " %buf_" + b0.name + " %c0u %i\n";
856 o += "OpStore %gep_out %final_val\n";
857
858 o += "OpReturn\n";
859 o += "OpFunctionEnd\n";
860 return o;
861 }
862
863 std::string emit_max_index_body(const ShaderSpec& spec)
864 {
865 const uint32_t local = spec.workgroup_size[0];
866 const std::string ls = std::to_string(local);
867 const auto& b0 = spec.bindings[0];
868 const auto& b1 = spec.bindings[1];
869
870 std::string o;
871 o += "%main = OpFunction %void None %voidfn\n";
872 o += "%entry = OpLabel\n";
873
874 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
875 o += "%i = OpCompositeExtract %u32 %gid3 0\n";
876 o += "%lid3 = OpLoad %v3u32 %lid_var\n";
877 o += "%lid = OpCompositeExtract %u32 %lid3 0\n\n";
878
879 o += "%gep_in = OpAccessChain %pelem_" + b0.name
880 + " %buf_" + b0.name + " %c0u %i\n";
881 o += "%elem = OpLoad %f32 %gep_in\n";
882 o += "%pgsh = OpAccessChain %psh_f32 %shared_a %lid\n";
883 o += "OpStore %pgsh %elem\n";
884 o += "%pgidx = OpAccessChain %psh_u32 %shared_idx %lid\n";
885 o += "OpStore %pgidx %lid\n";
886 o += "OpControlBarrier %c2u %c2u %c264u\n\n";
887
888 o += "%s_init = OpShiftRightLogical %u32 %c" + ls + "u %c1u\n";
889 o += "OpBranch %loop_hdr\n\n";
890
891 o += "%loop_hdr = OpLabel\n";
892 o += "%stride = OpPhi %u32 %s_init %entry %s_next %loop_cont\n";
893 o += "OpLoopMerge %loop_merge %loop_cont None\n";
894 o += "OpBranch %loop_body\n\n";
895
896 o += "%loop_body = OpLabel\n";
897 o += "%active = OpULessThan %bool %lid %stride\n";
898 o += "OpSelectionMerge %sel_merge None\n";
899 o += "OpBranchConditional %active %do_op %sel_merge\n\n";
900
901 o += "%do_op = OpLabel\n";
902 o += "%pgsh_a = OpAccessChain %psh_f32 %shared_a %lid\n";
903 o += "%a = OpLoad %f32 %pgsh_a\n";
904 o += "%lid_b = OpIAdd %u32 %lid %stride\n";
905 o += "%pgsh_b = OpAccessChain %psh_f32 %shared_a %lid_b\n";
906 o += "%b = OpLoad %f32 %pgsh_b\n";
907 o += "%pgidx_a = OpAccessChain %psh_u32 %shared_idx %lid\n";
908 o += "%idx_a = OpLoad %u32 %pgidx_a\n";
909 o += "%pgidx_b = OpAccessChain %psh_u32 %shared_idx %lid_b\n";
910 o += "%idx_b = OpLoad %u32 %pgidx_b\n";
911 o += "%b_wins = OpFOrdGreaterThan %bool %b %a\n";
912 o += "%combined = OpSelect %f32 %b_wins %b %a\n";
913 o += "%combined_idx = OpSelect %u32 %b_wins %idx_b %idx_a\n";
914 o += "OpStore %pgsh_a %combined\n";
915 o += "OpStore %pgidx_a %combined_idx\n";
916 o += "OpBranch %sel_merge\n\n";
917
918 o += "%sel_merge = OpLabel\n";
919 o += "OpBranch %loop_cont\n\n";
920
921 o += "%loop_cont = OpLabel\n";
922 o += "OpControlBarrier %c2u %c2u %c264u\n";
923 o += "%s_next = OpShiftRightLogical %u32 %stride %c1u\n";
924 o += "%done = OpIEqual %bool %s_next %c0u\n";
925 o += "OpBranchConditional %done %loop_merge %loop_hdr\n\n";
926
927 o += "%loop_merge = OpLabel\n";
928 o += "%is_zero = OpIEqual %bool %lid %c0u\n";
929 o += "OpSelectionMerge %write_merge None\n";
930 o += "OpBranchConditional %is_zero %do_write %write_merge\n\n";
931
932 o += "%do_write = OpLabel\n";
933 o += "%result_v = OpLoad %f32 %pgsh\n";
934 o += "%result_i = OpLoad %u32 %pgidx\n";
935 o += "%gep_out_v = OpAccessChain %pelem_" + b0.name
936 + " %buf_" + b0.name + " %c0u %c0u\n";
937 o += "OpStore %gep_out_v %result_v\n";
938 o += "%gep_out_i = OpAccessChain %pelem_" + b1.name
939 + " %buf_" + b1.name + " %c0u %c0u\n";
940 o += "OpStore %gep_out_i %result_i\n";
941 o += "OpBranch %write_merge\n\n";
942
943 o += "%write_merge = OpLabel\n";
944 o += "OpReturn\n";
945 o += "OpFunctionEnd\n";
946 return o;
947 }
948
949 /**
950 * Emit the entry point body for specs whose bindings include IMAGE_2D slots.
951 *
952 * Assumes workgroup_size is {8, 8, 1} or any 2D shape. GlobalInvocationId
953 * x/y components are used as the texel coordinate. One output IMAGE_2D
954 * binding is required. PC fields provide float operands identical to the
955 * SSBO elementwise path. The op is applied per-channel on the rgba vec4
956 * loaded from the first input IMAGE_2D, or on a zero vec4 if no input image
957 * is declared.
958 */
959 std::string emit_image_body(const ShaderSpec& spec)
960 {
961 const BindingSlot* img_out = nullptr;
962 std::vector<const BindingSlot*> img_inputs;
963 for (const auto& b : spec.bindings) {
964 if (b.modality != Kakshya::DataModality::IMAGE_2D)
965 continue;
966 if (b.direction == BindingDirection::Output && !img_out) {
967 img_out = &b;
968 } else if (b.direction != BindingDirection::Output) {
969 img_inputs.push_back(&b);
970 }
971 }
972
973 std::string o;
974 o += "%main = OpFunction %void None %voidfn\n";
975 o += "%entry = OpLabel\n";
976 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
977 o += "%ix = OpCompositeExtract %u32 %gid3 0\n";
978 o += "%iy = OpCompositeExtract %u32 %gid3 1\n";
979 o += "%six = OpBitcast %i32 %ix\n";
980 o += "%siy = OpBitcast %i32 %iy\n";
981 o += "%coord = OpCompositeConstruct %v2i32 %six %siy\n\n";
982
983 if (img_out)
984 o += "%img_out_val = OpLoad %img2d_t %img_" + std::string(img_out->name) + "\n";
985 for (size_t ii = 0; ii < img_inputs.size(); ++ii)
986 o += "%img_in_val" + std::to_string(ii) + " = OpLoad %img2d_t %img_" + img_inputs[ii]->name + "\n";
987 o += "\n";
988
989 for (size_t fi = 0; fi < spec.pc_fields.size(); ++fi) {
990 const auto& f = spec.pc_fields[fi];
991 o += "%ppc_" + f.name + " = OpAccessChain %ppc_f32 %pc %c"
992 + std::to_string(fi) + "u\n";
993 o += "%pc_" + f.name + " = OpLoad %f32 %ppc_" + f.name + "\n";
994 }
995 if (!spec.pc_fields.empty())
996 o += "\n";
997
998 std::vector<const BindingSlot*> tex_inputs;
999 for (const auto& b : spec.bindings) {
1000 if (b.modality == Kakshya::DataModality::TEXTURE_2D)
1001 tex_inputs.push_back(&b);
1002 }
1003
1004 if (!tex_inputs.empty()) {
1005 const std::string pw = spec.pc_fields.size() > 0
1006 ? ("%pc_" + spec.pc_fields[0].name)
1007 : "%lod_zero";
1008 const std::string ph = spec.pc_fields.size() > 1
1009 ? ("%pc_" + spec.pc_fields[1].name)
1010 : "%lod_zero";
1011 o += "%fix = OpConvertUToF %f32 %ix\n";
1012 o += "%fiy = OpConvertUToF %f32 %iy\n";
1013 o += "%u = OpFDiv %f32 %fix " + pw + "\n";
1014 o += "%v = OpFDiv %f32 %fiy " + ph + "\n";
1015 o += "%uv = OpCompositeConstruct %v2f32 %u %v\n\n";
1016
1017 for (size_t ti = 0; ti < tex_inputs.size(); ++ti) {
1018 const std::string idx = std::to_string(img_inputs.size() + ti);
1019 o += "%simg_" + tex_inputs[ti]->name
1020 + " = OpLoad %simgc_t %tex_" + tex_inputs[ti]->name + "\n";
1021 o += "%raw_in" + idx + " = OpImageSampleExplicitLod %v4f32 %simg_"
1022 + tex_inputs[ti]->name + " %uv Lod %lod_zero\n";
1023 }
1024 o += "\n";
1025 }
1026
1027 for (size_t ii = 0; ii < img_inputs.size(); ++ii) {
1028 const std::string idx = std::to_string(ii);
1029 o += "%raw_in" + idx + " = OpImageRead %v4f32 %img_in_val" + idx + " %coord\n";
1030 }
1031 if (img_inputs.empty()) {
1032 o += "%raw_in0 = OpCompositeConstruct %v4f32 %img_czero %img_czero"
1033 " %img_czero %img_czero\n";
1034 }
1035 o += "\n";
1036
1037 o += "%ch0_r = OpCompositeExtract %f32 %raw_in0 0\n";
1038 o += "%ch0_g = OpCompositeExtract %f32 %raw_in0 1\n";
1039 o += "%ch0_b = OpCompositeExtract %f32 %raw_in0 2\n";
1040 o += "%ch0_a = OpCompositeExtract %f32 %raw_in0 3\n";
1041
1042 const bool has_second = img_inputs.size() > 1
1043 || (!img_inputs.empty() && !tex_inputs.empty())
1044 || tex_inputs.size() > 1;
1045
1046 const std::string second_idx = img_inputs.size() > 1
1047 ? "1"
1048 : (!tex_inputs.empty() ? std::to_string(img_inputs.size()) : "");
1049
1050 if (has_second && !second_idx.empty()) {
1051 o += "%ch1_r = OpCompositeExtract %f32 %raw_in" + second_idx + " 0\n";
1052 o += "%ch1_g = OpCompositeExtract %f32 %raw_in" + second_idx + " 1\n";
1053 o += "%ch1_b = OpCompositeExtract %f32 %raw_in" + second_idx + " 2\n";
1054 o += "%ch1_a = OpCompositeExtract %f32 %raw_in" + second_idx + " 3\n";
1055 }
1056 o += "\n";
1057
1058 const std::string p0 = spec.pc_fields.empty()
1059 ? ""
1060 : ("%pc_" + spec.pc_fields[0].name);
1061 const std::string p1 = spec.pc_fields.size() > 1
1062 ? ("%pc_" + spec.pc_fields[1].name)
1063 : "";
1064
1065 if (spec.op == KernelOp::ChannelDot) {
1066 const std::string& wr = "%pc_" + spec.pc_fields[0].name;
1067 const std::string& wg = spec.pc_fields.size() > 1 ? "%pc_" + spec.pc_fields[1].name : "%img_czero";
1068 const std::string& wb = spec.pc_fields.size() > 2 ? "%pc_" + spec.pc_fields[2].name : "%img_czero";
1069 const std::string& wa = spec.pc_fields.size() > 3 ? "%pc_" + spec.pc_fields[3].name : "%img_czero";
1070 o += "%dot_r = OpFMul %f32 %ch0_r " + wr + "\n";
1071 o += "%dot_g = OpFMul %f32 %ch0_g " + wg + "\n";
1072 o += "%dot_b = OpFMul %f32 %ch0_b " + wb + "\n";
1073 o += "%dot_a = OpFMul %f32 %ch0_a " + wa + "\n";
1074 o += "%dot_rg = OpFAdd %f32 %dot_r %dot_g\n";
1075 o += "%dot_rgb = OpFAdd %f32 %dot_rg %dot_b\n";
1076 o += "%dot_val = OpFAdd %f32 %dot_rgb %dot_a\n";
1077 o += "%out_vec = OpCompositeConstruct %v4f32 %dot_val %dot_val %dot_val %dot_val\n";
1078 if (img_out)
1079 o += "OpImageWrite %img_out_val %coord %out_vec\n";
1080 o += "OpReturn\n";
1081 o += "OpFunctionEnd\n";
1082 return o;
1083 }
1084
1085 if (spec.op == KernelOp::ChannelReplicate) {
1086 o += "%out_vec = OpCompositeConstruct %v4f32 %ch0_r %ch0_r %ch0_r %img_cone\n";
1087 if (img_out)
1088 o += "OpImageWrite %img_out_val %coord %out_vec\n";
1089 o += "OpReturn\n";
1090 o += "OpFunctionEnd\n";
1091 return o;
1092 }
1093
1094 auto emit_channel_op = [&](
1095 const std::string& c0, const std::string& c1,
1096 const std::string& suffix) {
1097 switch (spec.op) {
1098 case KernelOp::Scale:
1099 o += "%res_" + suffix + " = OpFMul %f32 " + c0 + " " + p0 + "\n";
1100 break;
1102 o += "%mul_" + suffix + " = OpFMul %f32 " + c0 + " " + p0 + "\n";
1103 o += "%res_" + suffix + " = OpFAdd %f32 %mul_" + suffix + " " + p1 + "\n";
1104 break;
1105 case KernelOp::Offset:
1106 o += "%res_" + suffix + " = OpFAdd %f32 " + c0 + " " + p0 + "\n";
1107 break;
1108 case KernelOp::Clip:
1109 o += "%res_" + suffix + " = OpExtInst %f32 %glsl FClamp "
1110 + c0 + " " + p0 + " " + p1 + "\n";
1111 break;
1112 case KernelOp::Abs:
1113 o += "%res_" + suffix + " = OpExtInst %f32 %glsl FAbs " + c0 + "\n";
1114 break;
1115 case KernelOp::Negate:
1116 o += "%res_" + suffix + " = OpFNegate %f32 " + c0 + "\n";
1117 break;
1118 case KernelOp::Add:
1119 o += "%res_" + suffix + " = OpFAdd %f32 " + c0 + " " + c1 + "\n";
1120 break;
1121 case KernelOp::Multiply:
1122 o += "%res_" + suffix + " = OpFMul %f32 " + c0 + " " + c1 + "\n";
1123 break;
1124 case KernelOp::Mix:
1125 o += "%res_" + suffix + " = OpExtInst %f32 %glsl FMix "
1126 + c0 + " " + c1 + " " + p0 + "\n";
1127 break;
1128 case KernelOp::Sub:
1129 o += "%res_" + suffix + " = OpFSub %f32 " + c0 + " " + c1 + "\n";
1130 break;
1132 o += "%cmp_" + suffix + " = OpFOrdGreaterThanEqual %bool " + c0 + " " + p0 + "\n";
1133 o += "%res_" + suffix + " = OpSelect %f32 %cmp_" + suffix + " %img_cone %img_czero\n";
1134 break;
1136 o += "%cmp_" + suffix + " = OpFOrdGreaterThanEqual %bool " + c0 + " " + p0 + "\n";
1137 o += "%res_" + suffix + " = OpSelect %f32 %cmp_" + suffix + " " + p1 + " " + c0 + "\n";
1138 break;
1139 default:
1140 o += "%res_" + suffix + " = OpCopyObject %f32 " + c0 + "\n";
1141 break;
1142 }
1143 };
1144
1145 const std::string zero = "%img_czero";
1146 const std::string c1r = has_second ? "%ch1_r" : zero;
1147 const std::string c1g = has_second ? "%ch1_g" : zero;
1148 const std::string c1b = has_second ? "%ch1_b" : zero;
1149 const std::string c1a = has_second ? "%ch1_a" : zero;
1150
1151 emit_channel_op("%ch0_r", c1r, "r");
1152 emit_channel_op("%ch0_g", c1g, "g");
1153 emit_channel_op("%ch0_b", c1b, "b");
1154 emit_channel_op("%ch0_a", c1a, "a");
1155 o += "\n";
1156
1157 o += "%out_vec = OpCompositeConstruct %v4f32 %res_r %res_g %res_b %res_a\n";
1158 if (img_out)
1159 o += "OpImageWrite %img_out_val %coord %out_vec\n";
1160
1161 o += "OpReturn\n";
1162 o += "OpFunctionEnd\n";
1163 return o;
1164 }
1165
1166 /**
1167 * Emit the entry point body for BitonicSort template.
1168 * Loads the index, loads PC fields, loads SSBO elements, applies op, stores.
1169 */
1170 std::string emit_bitonic_body(const ShaderSpec& spec)
1171 {
1172 const auto& bkeys = spec.bindings[0];
1173 const auto& bidx = spec.bindings[1];
1174
1175 const std::string ktype(ssbo_elem_spirv_type(bkeys.format));
1176 const std::string itype(ssbo_elem_spirv_type(bidx.format));
1177
1178 std::string o;
1179 o += "%main = OpFunction %void None %voidfn\n";
1180 o += "%entry = OpLabel\n";
1181 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
1182 o += "%i = OpCompositeExtract %u32 %gid3 0\n\n";
1183
1184 o += "%ppc_stage = OpAccessChain %ppc_u32 %pc %c0u\n";
1185 o += "%stage = OpLoad %u32 %ppc_stage\n";
1186 o += "%ppc_pass = OpAccessChain %ppc_u32 %pc %c1u\n";
1187 o += "%pass = OpLoad %u32 %ppc_pass\n";
1188 o += "%ppc_count = OpAccessChain %ppc_u32 %pc %c2u\n";
1189 o += "%count = OpLoad %u32 %ppc_count\n";
1190 o += "%ppc_desc = OpAccessChain %ppc_u32 %pc %c3u\n";
1191 o += "%descending = OpLoad %u32 %ppc_desc\n\n";
1192
1193 o += "%c1u_shift = OpShiftLeftLogical %u32 %c1u %pass\n";
1194 o += "%partner = OpBitwiseXor %u32 %i %c1u_shift\n\n";
1195
1196 o += "%partner_le_i = OpULessThanEqual %bool %partner %i\n";
1197 o += "%i_oob = OpUGreaterThanEqual %bool %i %count\n";
1198 o += "%p_oob = OpUGreaterThanEqual %bool %partner %count\n";
1199 o += "%oob_raw = OpLogicalOr %bool %i_oob %p_oob\n";
1200 o += "%skip = OpLogicalOr %bool %partner_le_i %oob_raw\n\n";
1201
1202 o += "OpSelectionMerge %early_merge None\n";
1203 o += "OpBranchConditional %skip %early_ret %do_sort\n\n";
1204
1205 o += "%do_sort = OpLabel\n";
1206
1207 o += "%gep_ki = OpAccessChain %pelem_" + bkeys.name
1208 + " %buf_" + bkeys.name + " %c0u %i\n";
1209 o += "%key_i = OpLoad " + ktype + " %gep_ki\n";
1210 o += "%gep_kp = OpAccessChain %pelem_" + bkeys.name
1211 + " %buf_" + bkeys.name + " %c0u %partner\n";
1212 o += "%key_p = OpLoad " + ktype + " %gep_kp\n\n";
1213
1214 o += "%gep_ii = OpAccessChain %pelem_" + bidx.name
1215 + " %buf_" + bidx.name + " %c0u %i\n";
1216 o += "%idx_i = OpLoad " + itype + " %gep_ii\n";
1217 o += "%gep_ip = OpAccessChain %pelem_" + bidx.name
1218 + " %buf_" + bidx.name + " %c0u %partner\n";
1219 o += "%idx_p = OpLoad " + itype + " %gep_ip\n\n";
1220
1221 o += "%dir_shift = OpShiftRightLogical %u32 %i %stage\n";
1222 o += "%dir_bit = OpBitwiseAnd %u32 %dir_shift %c1u\n\n";
1223
1224 o += "%gt = OpFOrdGreaterThan %bool %key_i %key_p\n";
1225 o += "%gt_u = OpSelect %u32 %gt %c1u %c0u\n";
1226 o += "%xor1 = OpBitwiseXor %u32 %gt_u %dir_bit\n";
1227 o += "%xor2 = OpBitwiseXor %u32 %xor1 %descending\n";
1228 o += "%do_swap = OpINotEqual %bool %xor2 %c0u\n\n";
1229
1230 o += "%new_ki = OpSelect " + ktype + " %do_swap %key_p %key_i\n";
1231 o += "%new_kp = OpSelect " + ktype + " %do_swap %key_i %key_p\n";
1232 o += "%new_ii = OpSelect " + itype + " %do_swap %idx_p %idx_i\n";
1233 o += "%new_ip = OpSelect " + itype + " %do_swap %idx_i %idx_p\n\n";
1234
1235 o += "OpStore %gep_ki %new_ki\n";
1236 o += "OpStore %gep_kp %new_kp\n";
1237 o += "OpStore %gep_ii %new_ii\n";
1238 o += "OpStore %gep_ip %new_ip\n";
1239 o += "OpBranch %early_merge\n\n";
1240
1241 o += "%early_ret = OpLabel\n";
1242 o += "OpBranch %early_merge\n\n";
1243
1244 o += "%early_merge = OpLabel\n";
1245 o += "OpReturn\n";
1246 o += "OpFunctionEnd\n";
1247 return o;
1248 }
1249
1250 /**
1251 * Emit the entry point body for 2D convolution template.
1252 * Loads the thread coordinates, loads PC fields, loads kernel weights from SSBO,
1253 * applies convolution to the input image, stores to output image.
1254 */
1255 std::string emit_convolve2d_body(const ShaderSpec& spec)
1256 {
1257 const BindingSlot* img_out = nullptr;
1258 const BindingSlot* img_src = nullptr;
1259 const BindingSlot* kern_ssbo = nullptr;
1260 for (const auto& b : spec.bindings) {
1261 if (b.modality == Kakshya::DataModality::IMAGE_2D) {
1262 if (b.direction == BindingDirection::Output) {
1263 img_out = &b;
1264 } else {
1265 img_src = &b;
1266 }
1267 } else {
1268 kern_ssbo = &b;
1269 }
1270 }
1271
1272 std::string o;
1273 o += "%main = OpFunction %void None %voidfn\n";
1274 o += "%entry = OpLabel\n";
1275
1276 o += "%gid3 = OpLoad %v3u32 %gid_var\n";
1277 o += "%ix = OpCompositeExtract %u32 %gid3 0\n";
1278 o += "%iy = OpCompositeExtract %u32 %gid3 1\n";
1279
1280 o += "%ppc_radius = OpAccessChain %ppc_u32 %pc %c0u\n";
1281 o += "%pc_radius = OpLoad %u32 %ppc_radius\n";
1282 o += "%ppc_width = OpAccessChain %ppc_u32 %pc %c1u\n";
1283 o += "%pc_width = OpLoad %u32 %ppc_width\n";
1284 o += "%ppc_height = OpAccessChain %ppc_u32 %pc %c2u\n";
1285 o += "%pc_height = OpLoad %u32 %ppc_height\n";
1286
1287 o += "%oob_x = OpUGreaterThanEqual %bool %ix %pc_width\n";
1288 o += "%oob_y = OpUGreaterThanEqual %bool %iy %pc_height\n";
1289 o += "%oob = OpLogicalOr %bool %oob_x %oob_y\n";
1290 o += "OpSelectionMerge %main_merge None\n";
1291 o += "OpBranchConditional %oob %main_merge %conv_start\n\n";
1292
1293 o += "%conv_start = OpLabel\n";
1294
1295 o += "%img_src_val = OpLoad %img2d_t %img_" + std::string(img_src->name) + "\n";
1296 o += "%img_out_val = OpLoad %img2d_t %img_" + std::string(img_out->name) + "\n";
1297
1298 o += "%diam = OpIMul %u32 %pc_radius %c2u\n";
1299 o += "%diam1 = OpIAdd %u32 %diam %c1u\n";
1300
1301 o += "%six = OpBitcast %i32 %ix\n";
1302 o += "%siy = OpBitcast %i32 %iy\n";
1303 o += "%srad = OpBitcast %i32 %pc_radius\n";
1304 o += "%sw = OpBitcast %i32 %pc_width\n";
1305 o += "%sh = OpBitcast %i32 %pc_height\n";
1306 o += "%sw_1 = OpISub %i32 %sw %ci1\n";
1307 o += "%sh_1 = OpISub %i32 %sh %ci1\n";
1308
1309 o += "%czero4 = OpCompositeConstruct %v4f32 %img_czero %img_czero %img_czero %img_czero\n";
1310
1311 o += "OpBranch %ky_hdr\n\n";
1312
1313 o += "%ky_hdr = OpLabel\n";
1314 o += "%ky_u = OpPhi %u32 %c0u %conv_start %ky_next %ky_cont\n";
1315 o += "%acc_ky = OpPhi %v4f32 %czero4 %conv_start %acc_kx_done %ky_cont\n";
1316 o += "%ky_done = OpUGreaterThanEqual %bool %ky_u %diam1\n";
1317 o += "OpLoopMerge %ky_merge %ky_cont None\n";
1318 o += "OpBranchConditional %ky_done %ky_merge %kx_pre\n\n";
1319
1320 o += "%kx_pre = OpLabel\n";
1321 o += "%ky_si = OpBitcast %i32 %ky_u\n";
1322 o += "%ky_off = OpISub %i32 %ky_si %srad\n";
1323 o += "%sy_raw = OpIAdd %i32 %siy %ky_off\n";
1324 o += "%sy_lo = OpExtInst %i32 %glsl SMax %sy_raw %ci0\n";
1325 o += "%sy = OpExtInst %i32 %glsl SMin %sy_lo %sh_1\n";
1326 o += "OpBranch %kx_hdr\n\n";
1327
1328 o += "%kx_hdr = OpLabel\n";
1329 o += "%kx_u = OpPhi %u32 %c0u %kx_pre %kx_next %kx_cont\n";
1330 o += "%acc_kx = OpPhi %v4f32 %acc_ky %kx_pre %acc_new %kx_cont\n";
1331 o += "%kx_done = OpUGreaterThanEqual %bool %kx_u %diam1\n";
1332 o += "OpLoopMerge %kx_merge %kx_cont None\n";
1333 o += "OpBranchConditional %kx_done %kx_merge %kx_body\n\n";
1334
1335 o += "%kx_body = OpLabel\n";
1336 o += "%kx_si = OpBitcast %i32 %kx_u\n";
1337 o += "%kx_off = OpISub %i32 %kx_si %srad\n";
1338 o += "%sx_raw = OpIAdd %i32 %six %kx_off\n";
1339 o += "%sx_lo = OpExtInst %i32 %glsl SMax %sx_raw %ci0\n";
1340 o += "%sx = OpExtInst %i32 %glsl SMin %sx_lo %sw_1\n";
1341
1342 o += "%sc = OpCompositeConstruct %v2i32 %sx %sy\n";
1343 o += "%px = OpImageRead %v4f32 %img_src_val %sc\n";
1344
1345 o += "%kidx_r = OpIMul %u32 %ky_u %diam1\n";
1346 o += "%kidx = OpIAdd %u32 %kidx_r %kx_u\n";
1347 o += "%k_gep = OpAccessChain %pelem_" + std::string(kern_ssbo->name)
1348 + " %buf_" + std::string(kern_ssbo->name) + " %c0u %kidx\n";
1349 o += "%kw = OpLoad %f32 %k_gep\n";
1350
1351 o += "%kw4 = OpCompositeConstruct %v4f32 %kw %kw %kw %kw\n";
1352 o += "%prod = OpFMul %v4f32 %px %kw4\n";
1353 o += "%acc_new = OpFAdd %v4f32 %acc_kx %prod\n";
1354 o += "OpBranch %kx_cont\n\n";
1355
1356 o += "%kx_cont = OpLabel\n";
1357 o += "%kx_next = OpIAdd %u32 %kx_u %c1u\n";
1358 o += "OpBranch %kx_hdr\n\n";
1359
1360 o += "%kx_merge = OpLabel\n";
1361 o += "%acc_kx_done = OpPhi %v4f32 %acc_kx %kx_hdr\n";
1362 o += "OpBranch %ky_cont\n\n";
1363
1364 o += "%ky_cont = OpLabel\n";
1365 o += "%ky_next = OpIAdd %u32 %ky_u %c1u\n";
1366 o += "OpBranch %ky_hdr\n\n";
1367
1368 o += "%ky_merge = OpLabel\n";
1369 o += "%final_acc = OpPhi %v4f32 %acc_ky %ky_hdr\n";
1370
1371 o += "%out_coord = OpCompositeConstruct %v2i32 %six %siy\n";
1372 o += "OpImageWrite %img_out_val %out_coord %final_acc\n";
1373 o += "OpBranch %main_merge\n\n";
1374
1375 o += "%main_merge = OpLabel\n";
1376 o += "OpReturn\n";
1377 o += "OpFunctionEnd\n";
1378 return o;
1379 }
1380
1381} // namespace
1382
1383std::string emit_spirv_asm(const ShaderSpec& spec)
1384{
1385 std::string src;
1386 src += emit_header(spec);
1387 src += emit_decorations(spec);
1388 src += emit_types(spec);
1389
1390 if (spec.tmpl == KernelTemplate::Convolve2D) {
1391 src += emit_convolve2d_body(spec);
1392 return src;
1393 }
1394
1395 bool has_image = false;
1396 for (const auto& b : spec.bindings) {
1397 if (b.modality == Kakshya::DataModality::IMAGE_2D) {
1398 has_image = true;
1399 break;
1400 }
1401 }
1402
1403 if (has_image) {
1404 src += emit_image_body(spec);
1405 return src;
1406 }
1407
1408 switch (spec.tmpl) {
1410 src += (spec.op == KernelOp::MaxIndex) ? emit_max_index_body(spec) : emit_reduction_body(spec);
1411 break;
1413 src += emit_scan_body(spec);
1414 break;
1416 src += emit_bitonic_body(spec);
1417 break;
1421 default:
1422 src += emit_elementwise_body(spec);
1423 break;
1424 }
1425 return src;
1426}
1427
1428std::string emit_glsl_kernel(const ShaderSpec& spec)
1429{
1430 const auto& ws = spec.workgroup_size;
1431 const auto& ks = *spec.kernel;
1432
1433 bool has_image = false;
1434 for (const auto& b : spec.bindings) {
1435 if (b.modality == Kakshya::DataModality::IMAGE_2D)
1436 has_image = true;
1437 }
1438
1439 std::string o;
1440 o += "#version 460\n";
1441 o += "layout(local_size_x = " + std::to_string(ws[0])
1442 + ", local_size_y = " + std::to_string(ws[1])
1443 + ", local_size_z = " + std::to_string(ws[2]) + ") in;\n\n";
1444
1445 for (const auto& b : spec.bindings) {
1446 if (b.modality == Kakshya::DataModality::IMAGE_2D) {
1447 const std::string qual = (b.direction == BindingDirection::Input)
1448 ? "readonly"
1449 : "writeonly";
1450 o += "layout(set = 0, binding = " + std::to_string(b.binding_index)
1451 + ", rgba32f) " + qual + " uniform image2D " + b.name + ";\n";
1452 continue;
1453 }
1454 if (b.modality == Kakshya::DataModality::TEXTURE_2D) {
1455 o += "layout(set = 0, binding = " + std::to_string(b.binding_index)
1456 + ") uniform sampler2D " + b.name + ";\n";
1457 continue;
1458 }
1459 const auto t = std::string(glsl_type(b.format));
1460 o += "layout(set = 0, binding = " + std::to_string(b.binding_index)
1461 + ", std430) buffer Block_" + b.name
1462 + " { " + t + " " + b.name + "[]; };\n";
1463 }
1464
1465 if (!spec.pc_fields.empty()) {
1466 o += "\nlayout(push_constant) uniform PC {\n";
1467 for (const auto& f : spec.pc_fields)
1468 o += " " + std::string(glsl_type(f.format)) + " " + f.name + ";\n";
1469 o += "} pc;\n";
1470 }
1471
1472 o += "\nvoid main() {\n";
1473 o += " uint i = gl_GlobalInvocationID.x;\n";
1474 if (has_image)
1475 o += " ivec2 coord = ivec2(gl_GlobalInvocationID.xy);\n";
1476
1477 for (const auto& f : spec.pc_fields) {
1478 const auto t = std::string(glsl_type(f.format));
1479 o += " " + t + " " + f.name + " = pc." + f.name + ";\n";
1480 }
1481
1482 o += ks.body;
1483 o += "\n}\n";
1484 return o;
1485}
1486
1487} // namespace MayaFlux::Portal::Graphics::detail
size_t b
float wb
uint32_t pass
float wg
float wa
float wr
size_t gpu_data_format_bytes(GpuDataFormat fmt) noexcept
Byte size of one element of a GpuDataFormat.
Definition NDData.cpp:9
@ IMAGE_2D
2D image (grayscale or single channel)
GpuDataFormat
GPU data formats with explicit precision levels.
Definition NDData.hpp:25
std::string emit_glsl_kernel(const ShaderSpec &spec)
Emit a complete GLSL compute shader from spec metadata and a KernelSource body.
std::string emit_spirv_asm(const ShaderSpec &spec)
Emit complete SPIR-V assembly text for a generated compute kernel.
@ Scan
Inclusive prefix scan over one InOut SSBO, double-buffered Hillis-Steele in shared memory.
@ Convolve2D
2D separable or non-separable convolution; kernel weights in SSBO, radius in PC
@ Reduction
f(x[0..n]) -> scalar; shared-memory tree reduction
@ Elementwise
f(x[i]) -> y[i]; one thread per element
@ Stencil
f(x[i-k..i+k]) -> y[i]; neighbourhood reads, radius in PC
@ GeometryEmit
Writes into vertex SSBO with atomic counter.
@ BitonicSort
Bitonic sort network; one thread per element.
@ CompareGE
out[ch] = pixel[ch] >= pc[0] ? 1.0 : 0.0
@ ChannelDot
out = dot(pixel.rgba, pc[0..3]) broadcast to all channels
@ IndexScale
out[i] = float(i) * a[i].
@ CompareGEPreserve
out[ch] = pixel[ch] >= pc[0] ? pc[1] : pixel[ch]
@ MaxIndex
Reduction variant: finds max value AND its index.
@ ChannelReplicate
out = pixel[pc_channel_index].xxxx (single channel to all)
Declaration of one SSBO or image binding in a generated shader.
std::vector< PushConstantField > pc_fields
std::vector< BindingSlot > bindings
std::array< uint32_t, 3 > workgroup_size
std::optional< KernelSource > kernel
When set, KernelOp is ignored.
Complete declarative description of a generated compute shader.