108 std::string emit_header(
const ShaderSpec& spec)
112 bool has_storage_image =
false;
113 bool has_input_image =
false;
116 has_storage_image =
true;
118 has_input_image =
true;
122 std::string iface =
"%gid_var";
127 iface +=
" %tex_" +
b.name;
130 iface +=
" %buf_" +
b.name;
132 if (has_storage_image) {
135 iface +=
" %img_" +
b.name;
142 iface +=
" %lid_var";
145 o +=
"OpCapability Shader\n";
147 if (has_storage_image)
148 o +=
"OpCapability StorageImageWriteWithoutFormat\n";
150 o +=
"OpCapability StorageImageReadWithoutFormat\n";
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";
165 std::string emit_decorations(
const ShaderSpec& spec)
168 o +=
"OpDecorate %gid_var BuiltIn GlobalInvocationId\n";
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";
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";
195 o +=
"OpDecorate " + var +
" NonWritable\n";
197 o +=
"OpDecorate " + var +
" NonReadable\n";
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";
210 o +=
"OpDecorate %lid_var BuiltIn LocalInvocationId\n";
213 o +=
"OpDecorate %pc_blk Block\n";
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
>(
230 std::string emit_types(
const ShaderSpec& spec)
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";
242 bool need_v2f32 =
false;
243 bool need_v3f32 =
false;
244 bool need_v4f32 =
false;
266 o +=
"%v2f32 = OpTypeVector %f32 2\n";
268 o +=
"%v3f32 = OpTypeVector %f32 3\n";
270 bool need_i32 =
false;
279 bool has_image_2d =
false;
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))
290 if (need_i32 && !has_image_2d)
291 o +=
"%i32 = OpTypeInt 32 1\n";
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";
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";
321 o +=
"%img_" +
b.name +
" = OpVariable %ptr_img2d UniformConstant\n";
326 bool has_texture_2d =
false;
329 has_texture_2d =
true;
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";
344 o +=
"%tex_" +
b.name +
" = OpVariable %ptr_simg UniformConstant\n";
346 o +=
"%lod_zero = OpConstant %f32 0.0\n";
352 const std::string ls = std::to_string(local);
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";
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";
369 o +=
"%c" + ls +
"u = OpConstant %u32 " + ls +
"\n";
370 o +=
"%c2u = OpConstant %u32 2\n";
371 o +=
"%c264u = OpConstant %u32 264\n";
375 const std::string ls = std::to_string(local);
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";
391 bool pc_has_uint =
false;
392 bool pc_has_int =
false;
393 o +=
"%pc_blk = OpTypeStruct";
395 o +=
" " + std::string(ssbo_elem_spirv_type(f.format));
403 o +=
"%ppc = OpTypePointer PushConstant %pc_blk\n";
404 o +=
"%pc = OpVariable %ppc PushConstant\n";
405 o +=
"%ppc_f32 = OpTypePointer PushConstant %f32\n";
407 o +=
"%ppc_u32 = OpTypePointer PushConstant %u32\n";
409 o +=
"%ppc_i32 = OpTypePointer PushConstant %i32\n";
413 o +=
"%c0u = OpConstant %u32 0\n";
414 o +=
"%c1u = OpConstant %u32 1\n";
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";
429 std::string emit_elementwise_body(
const ShaderSpec& spec)
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";
437 for (
size_t fi = 0; fi < spec.
pc_fields.size(); ++fi) {
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";
450 std::vector<const BindingSlot*> ssbos;
459 for (
const auto*
b : ssbos) {
461 primary_fmt =
b->format;
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);
469 for (
const auto*
b : ssbos) {
470 o +=
"%gep_" +
b->name +
" = OpAccessChain %pelem_" +
b->name
471 +
" %buf_" +
b->name +
" %c0u %i\n";
473 o +=
"%val_" +
b->name +
" = OpLoad " + std::string(etype)
474 +
" %gep_" +
b->name +
"\n";
479 auto pc_operand = [&](
const std::string& field_name) -> std::string {
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)
488 o += construct +
"\n";
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() ?
""
496 const std::string p1 = spec.
pc_fields.size() > 1
501 const std::string et = std::string(etype);
505 o +=
"%res = OpVectorTimesScalar " + et +
" " + v0
506 +
" %pc_" + spec.
pc_fields[0].name +
"\n";
508 o +=
"%res = OpFMul " + et +
" " + v0 +
" " + p0 +
"\n";
515 o +=
"%scaled = OpVectorTimesScalar " + et +
" " + v0
516 +
" %pc_" + spec.
pc_fields[0].name +
"\n";
517 o +=
"%res = OpFAdd " + et +
" %scaled " + p1 +
"\n";
519 o +=
"%mul = OpFMul " + et +
" " + v0 +
" " + p0 +
"\n";
520 o +=
"%res = OpFAdd " + et +
" %mul " + p1 +
"\n";
526 o +=
"%res = OpFAdd " + et +
" " + v0 +
" " + p0 +
"\n";
530 o +=
"%res = OpExtInst " + et +
" %glsl FClamp " + v0
531 +
" " + p0 +
" " + p1 +
"\n";
535 o +=
"%res = OpExtInst " + et +
" %glsl FAbs " + v0 +
"\n";
539 o +=
"%res = OpFNegate " + et +
" " + v0 +
"\n";
543 o +=
"%res = OpFAdd " + et +
" " + v0 +
" " + v1 +
"\n";
547 o +=
"%res = OpFMul " + et +
" " + v0 +
" " + v1 +
"\n";
551 o +=
"%dlt = OpFSub %f32 " + v1 +
" " + v0 +
"\n";
552 o +=
"%scl = OpFMul %f32 %dlt " + p0 +
"\n";
553 o +=
"%res = OpFAdd %f32 " + v0 +
" %scl\n";
557 o +=
"%res = OpFSub " + et +
" " + v0 +
" " + v1 +
"\n";
561 o +=
"%res = OpExtInst %f32 %glsl Fma " + v0 +
" " + p0 +
" " + p1 +
"\n";
565 o +=
"%res = OpExtInst %f32 %glsl Floor " + v0 +
"\n";
569 o +=
"%res = OpExtInst %f32 %glsl Ceil " + v0 +
"\n";
573 o +=
"%res = OpExtInst %f32 %glsl Round " + v0 +
"\n";
577 o +=
"%res = OpExtInst %f32 %glsl Trunc " + v0 +
"\n";
581 o +=
"%res = OpExtInst %f32 %glsl Fract " + v0 +
"\n";
585 o +=
"%res = OpExtInst %f32 %glsl Sqrt " + v0 +
"\n";
589 o +=
"%res = OpExtInst %f32 %glsl InverseSqrt " + v0 +
"\n";
593 o +=
"%res = OpExtInst %f32 %glsl Sin " + v0 +
"\n";
597 o +=
"%res = OpExtInst %f32 %glsl Cos " + v0 +
"\n";
601 o +=
"%res = OpExtInst %f32 %glsl Tan " + v0 +
"\n";
605 o +=
"%res = OpExtInst %f32 %glsl Asin " + v0 +
"\n";
609 o +=
"%res = OpExtInst %f32 %glsl Acos " + v0 +
"\n";
613 o +=
"%res = OpExtInst %f32 %glsl Atan " + v0 +
"\n";
617 o +=
"%res = OpExtInst %f32 %glsl Sinh " + v0 +
"\n";
621 o +=
"%res = OpExtInst %f32 %glsl Cosh " + v0 +
"\n";
625 o +=
"%res = OpExtInst %f32 %glsl Tanh " + v0 +
"\n";
629 o +=
"%res = OpExtInst %f32 %glsl Exp " + v0 +
"\n";
633 o +=
"%res = OpExtInst %f32 %glsl Exp2 " + v0 +
"\n";
637 o +=
"%res = OpExtInst %f32 %glsl Log " + v0 +
"\n";
641 o +=
"%res = OpExtInst %f32 %glsl Log2 " + v0 +
"\n";
645 o +=
"%res = OpExtInst %f32 %glsl Pow " + v0 +
" " + v1 +
"\n";
649 o +=
"%res = OpExtInst %f32 %glsl Atan2 " + v0 +
" " + v1 +
"\n";
653 o +=
"%res = OpExtInst %f32 %glsl FMin " + v0 +
" " + v1 +
"\n";
657 o +=
"%res = OpExtInst %f32 %glsl FMax " + v0 +
" " + v1 +
"\n";
661 o +=
"%res = OpExtInst %f32 %glsl Step " + v0 +
" " + v1 +
"\n";
665 o +=
"%res = OpExtInst %f32 %glsl SmoothStep " + v0 +
" " + v1 +
" " + p0 +
"\n";
669 o +=
"%i_f32 = OpConvertUToF %f32 %i\n";
670 o +=
"%res = OpFMul " + et +
" " + v0 +
" %i_f32\n";
675 o +=
"%res = OpCopyObject " + et +
" " + v0 +
"\n";
681 for (
const auto*
b : ssbos) {
684 o +=
"OpStore %gep_" +
b->name +
" " + result +
"\n";
688 o +=
"OpFunctionEnd\n";
692 std::string emit_reduction_body(
const ShaderSpec& spec)
695 const std::string ls = std::to_string(local);
697 const auto& b0 = spec.
bindings.front();
700 o +=
"%main = OpFunction %void None %voidfn\n";
701 o +=
"%entry = OpLabel\n";
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";
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";
715 o +=
"%s_init = OpShiftRightLogical %u32 %c" + ls +
"u %c1u\n";
716 o +=
"OpBranch %loop_hdr\n\n";
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";
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";
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";
736 o +=
"%combined = OpExtInst %f32 %glsl FMax %a %b\n";
738 o +=
"%combined = OpFAdd %f32 %a %b\n";
741 o +=
"OpStore %pgsh_a %combined\n";
742 o +=
"OpBranch %sel_merge\n\n";
744 o +=
"%sel_merge = OpLabel\n";
745 o +=
"OpBranch %loop_cont\n\n";
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";
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";
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";
765 o +=
"%write_merge = OpLabel\n";
767 o +=
"OpFunctionEnd\n";
789 std::string emit_scan_body(
const ShaderSpec& spec)
792 const auto log2_local =
static_cast<uint32_t
>(std::log2(
static_cast<double>(local)));
793 const auto& b0 = spec.
bindings.front();
796 o +=
"%main = OpFunction %void None %voidfn\n";
797 o +=
"%entry = OpLabel\n";
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";
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";
811 std::string read_buf =
"%shared_a";
812 std::string write_buf =
"%shared_b";
813 std::string stride_val =
"%c1u";
816 const std::string ps = std::to_string(
pass);
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";
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";
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";
840 o +=
"%scan_merge_" + ps +
" = OpLabel\n";
841 o +=
"OpControlBarrier %c2u %c2u %c264u\n\n";
843 std::swap(read_buf, write_buf);
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;
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";
859 o +=
"OpFunctionEnd\n";
863 std::string emit_max_index_body(
const ShaderSpec& spec)
866 const std::string ls = std::to_string(local);
871 o +=
"%main = OpFunction %void None %voidfn\n";
872 o +=
"%entry = OpLabel\n";
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";
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";
888 o +=
"%s_init = OpShiftRightLogical %u32 %c" + ls +
"u %c1u\n";
889 o +=
"OpBranch %loop_hdr\n\n";
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";
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";
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";
918 o +=
"%sel_merge = OpLabel\n";
919 o +=
"OpBranch %loop_cont\n\n";
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";
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";
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";
943 o +=
"%write_merge = OpLabel\n";
945 o +=
"OpFunctionEnd\n";
959 std::string emit_image_body(
const ShaderSpec& spec)
962 std::vector<const BindingSlot*> img_inputs;
969 img_inputs.push_back(&
b);
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";
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";
989 for (
size_t fi = 0; fi < spec.
pc_fields.size(); ++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";
998 std::vector<const BindingSlot*> tex_inputs;
1001 tex_inputs.push_back(&
b);
1004 if (!tex_inputs.empty()) {
1005 const std::string pw = spec.
pc_fields.size() > 0
1008 const std::string ph = spec.
pc_fields.size() > 1
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";
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";
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";
1031 if (img_inputs.empty()) {
1032 o +=
"%raw_in0 = OpCompositeConstruct %v4f32 %img_czero %img_czero"
1033 " %img_czero %img_czero\n";
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";
1042 const bool has_second = img_inputs.size() > 1
1043 || (!img_inputs.empty() && !tex_inputs.empty())
1044 || tex_inputs.size() > 1;
1046 const std::string second_idx = img_inputs.size() > 1
1048 : (!tex_inputs.empty() ? std::to_string(img_inputs.size()) :
"");
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";
1058 const std::string p0 = spec.
pc_fields.empty()
1061 const std::string p1 = spec.
pc_fields.size() > 1
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";
1079 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1081 o +=
"OpFunctionEnd\n";
1086 o +=
"%out_vec = OpCompositeConstruct %v4f32 %ch0_r %ch0_r %ch0_r %img_cone\n";
1088 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1090 o +=
"OpFunctionEnd\n";
1094 auto emit_channel_op = [&](
1095 const std::string& c0,
const std::string& c1,
1096 const std::string& suffix) {
1099 o +=
"%res_" + suffix +
" = OpFMul %f32 " + c0 +
" " + p0 +
"\n";
1102 o +=
"%mul_" + suffix +
" = OpFMul %f32 " + c0 +
" " + p0 +
"\n";
1103 o +=
"%res_" + suffix +
" = OpFAdd %f32 %mul_" + suffix +
" " + p1 +
"\n";
1106 o +=
"%res_" + suffix +
" = OpFAdd %f32 " + c0 +
" " + p0 +
"\n";
1109 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FClamp "
1110 + c0 +
" " + p0 +
" " + p1 +
"\n";
1113 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FAbs " + c0 +
"\n";
1116 o +=
"%res_" + suffix +
" = OpFNegate %f32 " + c0 +
"\n";
1119 o +=
"%res_" + suffix +
" = OpFAdd %f32 " + c0 +
" " + c1 +
"\n";
1122 o +=
"%res_" + suffix +
" = OpFMul %f32 " + c0 +
" " + c1 +
"\n";
1125 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FMix "
1126 + c0 +
" " + c1 +
" " + p0 +
"\n";
1129 o +=
"%res_" + suffix +
" = OpFSub %f32 " + c0 +
" " + c1 +
"\n";
1132 o +=
"%cmp_" + suffix +
" = OpFOrdGreaterThanEqual %bool " + c0 +
" " + p0 +
"\n";
1133 o +=
"%res_" + suffix +
" = OpSelect %f32 %cmp_" + suffix +
" %img_cone %img_czero\n";
1136 o +=
"%cmp_" + suffix +
" = OpFOrdGreaterThanEqual %bool " + c0 +
" " + p0 +
"\n";
1137 o +=
"%res_" + suffix +
" = OpSelect %f32 %cmp_" + suffix +
" " + p1 +
" " + c0 +
"\n";
1140 o +=
"%res_" + suffix +
" = OpCopyObject %f32 " + c0 +
"\n";
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;
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");
1157 o +=
"%out_vec = OpCompositeConstruct %v4f32 %res_r %res_g %res_b %res_a\n";
1159 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1162 o +=
"OpFunctionEnd\n";
1170 std::string emit_bitonic_body(
const ShaderSpec& spec)
1172 const auto& bkeys = spec.
bindings[0];
1173 const auto& bidx = spec.
bindings[1];
1175 const std::string ktype(ssbo_elem_spirv_type(bkeys.format));
1176 const std::string itype(ssbo_elem_spirv_type(bidx.format));
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";
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";
1193 o +=
"%c1u_shift = OpShiftLeftLogical %u32 %c1u %pass\n";
1194 o +=
"%partner = OpBitwiseXor %u32 %i %c1u_shift\n\n";
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";
1202 o +=
"OpSelectionMerge %early_merge None\n";
1203 o +=
"OpBranchConditional %skip %early_ret %do_sort\n\n";
1205 o +=
"%do_sort = OpLabel\n";
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";
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";
1221 o +=
"%dir_shift = OpShiftRightLogical %u32 %i %stage\n";
1222 o +=
"%dir_bit = OpBitwiseAnd %u32 %dir_shift %c1u\n\n";
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";
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";
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";
1241 o +=
"%early_ret = OpLabel\n";
1242 o +=
"OpBranch %early_merge\n\n";
1244 o +=
"%early_merge = OpLabel\n";
1246 o +=
"OpFunctionEnd\n";
1255 std::string emit_convolve2d_body(
const ShaderSpec& spec)
1273 o +=
"%main = OpFunction %void None %voidfn\n";
1274 o +=
"%entry = OpLabel\n";
1276 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1277 o +=
"%ix = OpCompositeExtract %u32 %gid3 0\n";
1278 o +=
"%iy = OpCompositeExtract %u32 %gid3 1\n";
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";
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";
1293 o +=
"%conv_start = OpLabel\n";
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";
1298 o +=
"%diam = OpIMul %u32 %pc_radius %c2u\n";
1299 o +=
"%diam1 = OpIAdd %u32 %diam %c1u\n";
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";
1309 o +=
"%czero4 = OpCompositeConstruct %v4f32 %img_czero %img_czero %img_czero %img_czero\n";
1311 o +=
"OpBranch %ky_hdr\n\n";
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";
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";
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";
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";
1342 o +=
"%sc = OpCompositeConstruct %v2i32 %sx %sy\n";
1343 o +=
"%px = OpImageRead %v4f32 %img_src_val %sc\n";
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";
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";
1356 o +=
"%kx_cont = OpLabel\n";
1357 o +=
"%kx_next = OpIAdd %u32 %kx_u %c1u\n";
1358 o +=
"OpBranch %kx_hdr\n\n";
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";
1364 o +=
"%ky_cont = OpLabel\n";
1365 o +=
"%ky_next = OpIAdd %u32 %ky_u %c1u\n";
1366 o +=
"OpBranch %ky_hdr\n\n";
1368 o +=
"%ky_merge = OpLabel\n";
1369 o +=
"%final_acc = OpPhi %v4f32 %acc_ky %ky_hdr\n";
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";
1375 o +=
"%main_merge = OpLabel\n";
1377 o +=
"OpFunctionEnd\n";
1386 src += emit_header(spec);
1387 src += emit_decorations(spec);
1388 src += emit_types(spec);
1391 src += emit_convolve2d_body(spec);
1395 bool has_image =
false;
1404 src += emit_image_body(spec);
1408 switch (spec.
tmpl) {
1410 src += (spec.
op ==
KernelOp::MaxIndex) ? emit_max_index_body(spec) : emit_reduction_body(spec);
1413 src += emit_scan_body(spec);
1416 src += emit_bitonic_body(spec);
1422 src += emit_elementwise_body(spec);
1431 const auto& ks = *spec.
kernel;
1433 bool has_image =
false;
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";
1450 o +=
"layout(set = 0, binding = " + std::to_string(
b.binding_index)
1451 +
", rgba32f) " + qual +
" uniform image2D " +
b.name +
";\n";
1455 o +=
"layout(set = 0, binding = " + std::to_string(
b.binding_index)
1456 +
") uniform sampler2D " +
b.name +
";\n";
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";
1466 o +=
"\nlayout(push_constant) uniform PC {\n";
1468 o +=
" " + std::string(glsl_type(f.format)) +
" " + f.name +
";\n";
1472 o +=
"\nvoid main() {\n";
1473 o +=
" uint i = gl_GlobalInvocationID.x;\n";
1475 o +=
" ivec2 coord = ivec2(gl_GlobalInvocationID.xy);\n";
1478 const auto t = std::string(glsl_type(f.format));
1479 o +=
" " + t +
" " + f.name +
" = pc." + f.name +
";\n";