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";
443 std::string emit_stencil_body(
const ShaderSpec& spec)
455 const std::string_view etype = ssbo_elem_spirv_type(in_ssbo->
format);
456 const bool is_float = (etype ==
"%f32");
457 const std::string add_op = is_float ?
"OpFAdd" :
"OpIAdd";
458 const std::string sub_op = is_float ?
"OpFSub" :
"OpISub";
459 const std::string mul_op = is_float ?
"OpFMul" :
"OpIMul";
460 const std::string in_name = in_ssbo->
name;
461 const std::string out_name = out_ssbo->
name;
464 o +=
"%main = OpFunction %void None %voidfn\n";
465 o +=
"%entry = OpLabel\n";
466 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
467 o +=
"%ix = OpCompositeExtract %u32 %gid3 0\n";
468 o +=
"%iy = OpCompositeExtract %u32 %gid3 1\n\n";
470 o +=
"%ppc_width = OpAccessChain %ppc_u32 %pc %c0u\n";
471 o +=
"%pc_width = OpLoad %u32 %ppc_width\n";
472 o +=
"%ppc_height = OpAccessChain %ppc_u32 %pc %c1u\n";
473 o +=
"%pc_height = OpLoad %u32 %ppc_height\n\n";
475 o +=
"%w_m1 = OpISub %u32 %pc_width %c1u\n";
476 o +=
"%h_m1 = OpISub %u32 %pc_height %c1u\n\n";
478 o +=
"%i = OpIMul %u32 %iy %pc_width\n";
479 o +=
"%i_flat = OpIAdd %u32 %i %ix\n\n";
481 const std::array<std::pair<int, int>, 9> offsets = { { { -1, -1 }, { 0, -1 }, { 1, -1 },
482 { -1, 0 }, { 0, 0 }, { 1, 0 },
483 { -1, 1 }, { 0, 1 }, { 1, 1 } } };
485 std::string center_val;
486 std::string running_sum;
487 bool sum_started =
false;
489 for (
size_t n = 0; n < offsets.size(); ++n) {
490 const auto [dx, dy] = offsets[n];
491 const std::string ns = std::to_string(n);
493 std::string nx =
"%ix";
495 o +=
"%nxr" + ns +
" = OpIAdd %u32 %ix %c1u\n";
496 o +=
"%oobx" + ns +
" = OpUGreaterThan %bool %nxr" + ns +
" %w_m1\n";
497 o +=
"%nx" + ns +
" = OpSelect %u32 %oobx" + ns +
" %w_m1 %nxr" + ns +
"\n";
499 }
else if (dx == -1) {
500 o +=
"%z" + ns +
" = OpIEqual %bool %ix %c0u\n";
501 o +=
"%nxr" + ns +
" = OpISub %u32 %ix %c1u\n";
502 o +=
"%nx" + ns +
" = OpSelect %u32 %z" + ns +
" %c0u %nxr" + ns +
"\n";
506 std::string ny =
"%iy";
508 o +=
"%nyr" + ns +
" = OpIAdd %u32 %iy %c1u\n";
509 o +=
"%ooby" + ns +
" = OpUGreaterThan %bool %nyr" + ns +
" %h_m1\n";
510 o +=
"%ny" + ns +
" = OpSelect %u32 %ooby" + ns +
" %h_m1 %nyr" + ns +
"\n";
512 }
else if (dy == -1) {
513 o +=
"%zy" + ns +
" = OpIEqual %bool %iy %c0u\n";
514 o +=
"%nyr" + ns +
" = OpISub %u32 %iy %c1u\n";
515 o +=
"%ny" + ns +
" = OpSelect %u32 %zy" + ns +
" %c0u %nyr" + ns +
"\n";
519 o +=
"%nrow" + ns +
" = OpIMul %u32 " + ny +
" %pc_width\n";
520 o +=
"%nidx" + ns +
" = OpIAdd %u32 %nrow" + ns +
" " + nx +
"\n";
521 o +=
"%ngep" + ns +
" = OpAccessChain %pelem_" + in_name
522 +
" %buf_" + in_name +
" %c0u %nidx" + ns +
"\n";
523 o +=
"%nval" + ns +
" = OpLoad " + std::string(etype) +
" %ngep" + ns +
"\n";
525 if (dx == 0 && dy == 0) {
526 center_val =
"%nval" + ns;
531 running_sum =
"%nval" + ns;
534 const std::string next_sum =
"%sum" + ns;
535 o += next_sum +
" = " + add_op +
" " + std::string(etype) +
" " + running_sum +
" %nval" + ns +
"\n";
536 running_sum = next_sum;
544 o +=
"%res = " + add_op +
" " + std::string(etype) +
" " + center_val +
" " + running_sum +
"\n";
548 o +=
"%res = " + sub_op +
" " + std::string(etype) +
" " + center_val +
" " + running_sum +
"\n";
552 o +=
"%res = " + mul_op +
" " + std::string(etype) +
" " + center_val +
" " + running_sum +
"\n";
556 o +=
"%ppc_rate = OpAccessChain %ppc_f32 %pc %c2u\n";
557 o +=
"%pc_rate = OpLoad %f32 %ppc_rate\n";
558 o +=
"%ppc_wsum = OpAccessChain %ppc_f32 %pc %c3u\n";
559 o +=
"%pc_wsum = OpLoad %f32 %ppc_wsum\n";
560 o +=
"%scaled_sum = OpFMul %f32 " + running_sum +
" %pc_wsum\n";
561 o +=
"%delta = OpFSub %f32 %scaled_sum " + center_val +
"\n";
562 o +=
"%weighted = OpFMul %f32 %pc_rate %delta\n";
563 o +=
"%res = OpFAdd %f32 " + center_val +
" %weighted\n";
567 result = running_sum;
572 o +=
"%out_gep = OpAccessChain %pelem_" + out_name
573 +
" %buf_" + out_name +
" %c0u %i_flat\n";
574 o +=
"OpStore %out_gep " + result +
"\n";
577 o +=
"OpFunctionEnd\n";
585 std::string emit_elementwise_body(
const ShaderSpec& spec)
588 o +=
"%main = OpFunction %void None %voidfn\n";
589 o +=
"%entry = OpLabel\n";
590 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
591 o +=
"%i = OpCompositeExtract %u32 %gid3 0\n\n";
593 for (
size_t fi = 0; fi < spec.
pc_fields.size(); ++fi) {
598 const std::string_view ld_type = ssbo_elem_spirv_type(f.format);
599 o +=
"%ppc_" + f.name +
" = OpAccessChain " + std::string(pptr)
600 +
" %pc %c" + std::to_string(fi) +
"u\n";
601 o +=
"%pc_" + f.name +
" = OpLoad " + std::string(ld_type)
602 +
" %ppc_" + f.name +
"\n";
606 std::vector<const BindingSlot*> ssbos;
615 for (
const auto*
b : ssbos) {
617 primary_fmt =
b->format;
621 const std::string_view etype = ssbo_elem_spirv_type(primary_fmt);
622 const uint32_t ncomp = ssbo_elem_components(primary_fmt);
623 const bool is_vector = (ncomp > 1);
625 for (
const auto*
b : ssbos) {
626 o +=
"%gep_" +
b->name +
" = OpAccessChain %pelem_" +
b->name
627 +
" %buf_" +
b->name +
" %c0u %i\n";
629 o +=
"%val_" +
b->name +
" = OpLoad " + std::string(etype)
630 +
" %gep_" +
b->name +
"\n";
635 auto pc_operand = [&](
const std::string& field_name) -> std::string {
637 return "%pc_" + field_name;
638 const std::string splat =
"%spc_" + field_name;
639 std::string construct = splat +
" = OpCompositeConstruct "
640 + std::string(etype);
641 const std::string scalar =
" %pc_" + field_name;
642 for (uint32_t c = 0; c < ncomp; ++c)
644 o += construct +
"\n";
648 const std::string v0 = ssbos.empty() ?
"" : (
"%val_" + ssbos[0]->name);
649 const std::string v1 = ssbos.size() > 1 ? (
"%val_" + ssbos[1]->name) :
"";
650 const std::string p0 = spec.
pc_fields.empty() ?
""
652 const std::string p1 = spec.
pc_fields.size() > 1
657 const std::string et = std::string(etype);
661 o +=
"%res = OpVectorTimesScalar " + et +
" " + v0
662 +
" %pc_" + spec.
pc_fields[0].name +
"\n";
664 o +=
"%res = OpFMul " + et +
" " + v0 +
" " + p0 +
"\n";
671 o +=
"%scaled = OpVectorTimesScalar " + et +
" " + v0
672 +
" %pc_" + spec.
pc_fields[0].name +
"\n";
673 o +=
"%res = OpFAdd " + et +
" %scaled " + p1 +
"\n";
675 o +=
"%mul = OpFMul " + et +
" " + v0 +
" " + p0 +
"\n";
676 o +=
"%res = OpFAdd " + et +
" %mul " + p1 +
"\n";
682 o +=
"%res = OpFAdd " + et +
" " + v0 +
" " + p0 +
"\n";
686 o +=
"%res = OpExtInst " + et +
" %glsl FClamp " + v0
687 +
" " + p0 +
" " + p1 +
"\n";
691 o +=
"%res = OpExtInst " + et +
" %glsl FAbs " + v0 +
"\n";
695 o +=
"%res = OpFNegate " + et +
" " + v0 +
"\n";
699 o +=
"%res = OpFAdd " + et +
" " + v0 +
" " + v1 +
"\n";
703 o +=
"%res = OpFMul " + et +
" " + v0 +
" " + v1 +
"\n";
707 o +=
"%dlt = OpFSub %f32 " + v1 +
" " + v0 +
"\n";
708 o +=
"%scl = OpFMul %f32 %dlt " + p0 +
"\n";
709 o +=
"%res = OpFAdd %f32 " + v0 +
" %scl\n";
713 o +=
"%res = OpFSub " + et +
" " + v0 +
" " + v1 +
"\n";
717 o +=
"%res = OpExtInst %f32 %glsl Fma " + v0 +
" " + p0 +
" " + p1 +
"\n";
721 o +=
"%res = OpExtInst %f32 %glsl Floor " + v0 +
"\n";
725 o +=
"%res = OpExtInst %f32 %glsl Ceil " + v0 +
"\n";
729 o +=
"%res = OpExtInst %f32 %glsl Round " + v0 +
"\n";
733 o +=
"%res = OpExtInst %f32 %glsl Trunc " + v0 +
"\n";
737 o +=
"%res = OpExtInst %f32 %glsl Fract " + v0 +
"\n";
741 o +=
"%res = OpExtInst %f32 %glsl Sqrt " + v0 +
"\n";
745 o +=
"%res = OpExtInst %f32 %glsl InverseSqrt " + v0 +
"\n";
749 o +=
"%res = OpExtInst %f32 %glsl Sin " + v0 +
"\n";
753 o +=
"%res = OpExtInst %f32 %glsl Cos " + v0 +
"\n";
757 o +=
"%res = OpExtInst %f32 %glsl Tan " + v0 +
"\n";
761 o +=
"%res = OpExtInst %f32 %glsl Asin " + v0 +
"\n";
765 o +=
"%res = OpExtInst %f32 %glsl Acos " + v0 +
"\n";
769 o +=
"%res = OpExtInst %f32 %glsl Atan " + v0 +
"\n";
773 o +=
"%res = OpExtInst %f32 %glsl Sinh " + v0 +
"\n";
777 o +=
"%res = OpExtInst %f32 %glsl Cosh " + v0 +
"\n";
781 o +=
"%res = OpExtInst %f32 %glsl Tanh " + v0 +
"\n";
785 o +=
"%res = OpExtInst %f32 %glsl Exp " + v0 +
"\n";
789 o +=
"%res = OpExtInst %f32 %glsl Exp2 " + v0 +
"\n";
793 o +=
"%res = OpExtInst %f32 %glsl Log " + v0 +
"\n";
797 o +=
"%res = OpExtInst %f32 %glsl Log2 " + v0 +
"\n";
801 o +=
"%res = OpExtInst %f32 %glsl Pow " + v0 +
" " + v1 +
"\n";
805 o +=
"%res = OpExtInst %f32 %glsl Atan2 " + v0 +
" " + v1 +
"\n";
809 o +=
"%res = OpExtInst %f32 %glsl FMin " + v0 +
" " + v1 +
"\n";
813 o +=
"%res = OpExtInst %f32 %glsl FMax " + v0 +
" " + v1 +
"\n";
817 o +=
"%res = OpExtInst %f32 %glsl Step " + v0 +
" " + v1 +
"\n";
821 o +=
"%res = OpExtInst %f32 %glsl SmoothStep " + v0 +
" " + v1 +
" " + p0 +
"\n";
825 o +=
"%i_f32 = OpConvertUToF %f32 %i\n";
826 o +=
"%res = OpFMul " + et +
" " + v0 +
" %i_f32\n";
832 uint32_t out_ncomp = ncomp;
833 for (
const auto*
b : ssbos) {
836 out_ncomp = ssbo_elem_components(out_fmt);
840 const std::string_view out_etype = ssbo_elem_spirv_type(out_fmt);
842 std::string cmp_lhs = v0;
844 o +=
"%cmp_scalar = OpCompositeExtract %f32 " + v0 +
" 0\n";
845 cmp_lhs =
"%cmp_scalar";
846 }
else if (etype !=
"%f32") {
847 o +=
"%cmp_lhs_f = OpConvertUToF %f32 " + v0 +
"\n";
848 cmp_lhs =
"%cmp_lhs_f";
851 const bool has_threshold = !spec.
pc_fields.empty();
852 std::string threshold_f;
856 o +=
"%cmp_zero_u = OpConvertUToF %f32 %c0u\n";
857 threshold_f =
"%cmp_zero_u";
860 o +=
"%cmp = OpFOrdGreaterThanEqual %bool " + cmp_lhs +
" " + threshold_f +
"\n";
861 o +=
"%cmp_one_f = OpConvertUToF %f32 %c1u\n";
862 o +=
"%cmp_zero_f = OpConvertUToF %f32 %c0u\n";
863 o +=
"%cmp_scaled = OpSelect %f32 %cmp %cmp_one_f %cmp_zero_f\n";
866 o +=
"%res = OpCompositeConstruct " + std::string(out_etype);
867 for (uint32_t c = 0; c < out_ncomp; ++c)
870 }
else if (out_etype !=
"%f32") {
871 o +=
"%res = OpConvertFToU " + std::string(out_etype) +
" %cmp_scaled\n";
873 o +=
"%res = OpCopyObject %f32 %cmp_scaled\n";
879 for (
const auto*
b : ssbos) {
883 o +=
"%res = OpCopyObject " + et +
" " + v0 +
"\n";
889 for (
const auto*
b : ssbos) {
892 o +=
"OpStore %gep_" +
b->name +
" " + result +
"\n";
896 o +=
"OpFunctionEnd\n";
900 std::string emit_reduction_body(
const ShaderSpec& spec)
903 const std::string ls = std::to_string(local);
905 const auto& b0 = spec.
bindings.front();
908 o +=
"%main = OpFunction %void None %voidfn\n";
909 o +=
"%entry = OpLabel\n";
911 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
912 o +=
"%i = OpCompositeExtract %u32 %gid3 0\n";
913 o +=
"%lid3 = OpLoad %v3u32 %lid_var\n";
914 o +=
"%lid = OpCompositeExtract %u32 %lid3 0\n\n";
916 o +=
"%gep_in = OpAccessChain %pelem_" + b0.name
917 +
" %buf_" + b0.name +
" %c0u %i\n";
918 o +=
"%elem = OpLoad %f32 %gep_in\n";
919 o +=
"%pgsh = OpAccessChain %psh_f32 %shared_a %lid\n";
920 o +=
"OpStore %pgsh %elem\n";
921 o +=
"OpControlBarrier %c2u %c2u %c264u\n\n";
923 o +=
"%s_init = OpShiftRightLogical %u32 %c" + ls +
"u %c1u\n";
924 o +=
"OpBranch %loop_hdr\n\n";
926 o +=
"%loop_hdr = OpLabel\n";
927 o +=
"%stride = OpPhi %u32 %s_init %entry %s_next %loop_cont\n";
928 o +=
"OpLoopMerge %loop_merge %loop_cont None\n";
929 o +=
"OpBranch %loop_body\n\n";
931 o +=
"%loop_body = OpLabel\n";
932 o +=
"%active = OpULessThan %bool %lid %stride\n";
933 o +=
"OpSelectionMerge %sel_merge None\n";
934 o +=
"OpBranchConditional %active %do_op %sel_merge\n\n";
936 o +=
"%do_op = OpLabel\n";
937 o +=
"%pgsh_a = OpAccessChain %psh_f32 %shared_a %lid\n";
938 o +=
"%a = OpLoad %f32 %pgsh_a\n";
939 o +=
"%lid_b = OpIAdd %u32 %lid %stride\n";
940 o +=
"%pgsh_b = OpAccessChain %psh_f32 %shared_a %lid_b\n";
941 o +=
"%b = OpLoad %f32 %pgsh_b\n";
944 o +=
"%combined = OpExtInst %f32 %glsl FMax %a %b\n";
946 o +=
"%combined = OpFAdd %f32 %a %b\n";
949 o +=
"OpStore %pgsh_a %combined\n";
950 o +=
"OpBranch %sel_merge\n\n";
952 o +=
"%sel_merge = OpLabel\n";
953 o +=
"OpBranch %loop_cont\n\n";
955 o +=
"%loop_cont = OpLabel\n";
956 o +=
"OpControlBarrier %c2u %c2u %c264u\n";
957 o +=
"%s_next = OpShiftRightLogical %u32 %stride %c1u\n";
958 o +=
"%done = OpIEqual %bool %s_next %c0u\n";
959 o +=
"OpBranchConditional %done %loop_merge %loop_hdr\n\n";
961 o +=
"%loop_merge = OpLabel\n";
962 o +=
"%is_zero = OpIEqual %bool %lid %c0u\n";
963 o +=
"OpSelectionMerge %write_merge None\n";
964 o +=
"OpBranchConditional %is_zero %do_write %write_merge\n\n";
966 o +=
"%do_write = OpLabel\n";
967 o +=
"%result = OpLoad %f32 %pgsh\n";
968 o +=
"%gep_out = OpAccessChain %pelem_" + b0.name
969 +
" %buf_" + b0.name +
" %c0u %c0u\n";
970 o +=
"OpStore %gep_out %result\n";
971 o +=
"OpBranch %write_merge\n\n";
973 o +=
"%write_merge = OpLabel\n";
975 o +=
"OpFunctionEnd\n";
997 std::string emit_scan_body(
const ShaderSpec& spec)
1000 const auto log2_local =
static_cast<uint32_t
>(std::log2(
static_cast<double>(local)));
1001 const auto& b0 = spec.
bindings.front();
1004 o +=
"%main = OpFunction %void None %voidfn\n";
1005 o +=
"%entry = OpLabel\n";
1007 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1008 o +=
"%i = OpCompositeExtract %u32 %gid3 0\n";
1009 o +=
"%lid3 = OpLoad %v3u32 %lid_var\n";
1010 o +=
"%lid = OpCompositeExtract %u32 %lid3 0\n\n";
1012 o +=
"%gep_in = OpAccessChain %pelem_" + b0.name
1013 +
" %buf_" + b0.name +
" %c0u %i\n";
1014 o +=
"%elem = OpLoad %f32 %gep_in\n";
1015 o +=
"%pgsh0 = OpAccessChain %psh_f32 %shared_a %lid\n";
1016 o +=
"OpStore %pgsh0 %elem\n";
1017 o +=
"OpControlBarrier %c2u %c2u %c264u\n\n";
1019 std::string read_buf =
"%shared_a";
1020 std::string write_buf =
"%shared_b";
1021 std::string stride_val =
"%c1u";
1024 const std::string ps = std::to_string(
pass);
1026 o +=
"%has_left_" + ps +
" = OpUGreaterThanEqual %bool %lid " + stride_val +
"\n";
1027 o +=
"OpSelectionMerge %scan_merge_" + ps +
" None\n";
1028 o +=
"OpBranchConditional %has_left_" + ps +
" %scan_add_" + ps +
" %scan_copy_" + ps +
"\n\n";
1030 o +=
"%scan_add_" + ps +
" = OpLabel\n";
1031 o +=
"%self_ptr_" + ps +
" = OpAccessChain %psh_f32 " + read_buf +
" %lid\n";
1032 o +=
"%self_val_" + ps +
" = OpLoad %f32 %self_ptr_" + ps +
"\n";
1033 o +=
"%left_idx_" + ps +
" = OpISub %u32 %lid " + stride_val +
"\n";
1034 o +=
"%left_ptr_" + ps +
" = OpAccessChain %psh_f32 " + read_buf +
" %left_idx_" + ps +
"\n";
1035 o +=
"%left_val_" + ps +
" = OpLoad %f32 %left_ptr_" + ps +
"\n";
1036 o +=
"%sum_" + ps +
" = OpFAdd %f32 %self_val_" + ps +
" %left_val_" + ps +
"\n";
1037 o +=
"%wptr_add_" + ps +
" = OpAccessChain %psh_f32 " + write_buf +
" %lid\n";
1038 o +=
"OpStore %wptr_add_" + ps +
" %sum_" + ps +
"\n";
1039 o +=
"OpBranch %scan_merge_" + ps +
"\n\n";
1041 o +=
"%scan_copy_" + ps +
" = OpLabel\n";
1042 o +=
"%pass_ptr_" + ps +
" = OpAccessChain %psh_f32 " + read_buf +
" %lid\n";
1043 o +=
"%pass_val_" + ps +
" = OpLoad %f32 %pass_ptr_" + ps +
"\n";
1044 o +=
"%wptr_copy_" + ps +
" = OpAccessChain %psh_f32 " + write_buf +
" %lid\n";
1045 o +=
"OpStore %wptr_copy_" + ps +
" %pass_val_" + ps +
"\n";
1046 o +=
"OpBranch %scan_merge_" + ps +
"\n\n";
1048 o +=
"%scan_merge_" + ps +
" = OpLabel\n";
1049 o +=
"OpControlBarrier %c2u %c2u %c264u\n\n";
1051 std::swap(read_buf, write_buf);
1053 if (
pass + 1 < log2_local) {
1054 const std::string next_stride =
"%stride_next_" + ps;
1055 o += next_stride +
" = OpShiftLeftLogical %u32 " + stride_val +
" %c1u\n";
1056 stride_val = next_stride;
1060 o +=
"%final_ptr = OpAccessChain %psh_f32 " + read_buf +
" %lid\n";
1061 o +=
"%final_val = OpLoad %f32 %final_ptr\n";
1062 o +=
"%gep_out = OpAccessChain %pelem_" + b0.name
1063 +
" %buf_" + b0.name +
" %c0u %i\n";
1064 o +=
"OpStore %gep_out %final_val\n";
1067 o +=
"OpFunctionEnd\n";
1071 std::string emit_max_index_body(
const ShaderSpec& spec)
1074 const std::string ls = std::to_string(local);
1079 o +=
"%main = OpFunction %void None %voidfn\n";
1080 o +=
"%entry = OpLabel\n";
1082 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1083 o +=
"%i = OpCompositeExtract %u32 %gid3 0\n";
1084 o +=
"%lid3 = OpLoad %v3u32 %lid_var\n";
1085 o +=
"%lid = OpCompositeExtract %u32 %lid3 0\n\n";
1087 o +=
"%gep_in = OpAccessChain %pelem_" + b0.name
1088 +
" %buf_" + b0.name +
" %c0u %i\n";
1089 o +=
"%elem = OpLoad %f32 %gep_in\n";
1090 o +=
"%pgsh = OpAccessChain %psh_f32 %shared_a %lid\n";
1091 o +=
"OpStore %pgsh %elem\n";
1092 o +=
"%pgidx = OpAccessChain %psh_u32 %shared_idx %lid\n";
1093 o +=
"OpStore %pgidx %lid\n";
1094 o +=
"OpControlBarrier %c2u %c2u %c264u\n\n";
1096 o +=
"%s_init = OpShiftRightLogical %u32 %c" + ls +
"u %c1u\n";
1097 o +=
"OpBranch %loop_hdr\n\n";
1099 o +=
"%loop_hdr = OpLabel\n";
1100 o +=
"%stride = OpPhi %u32 %s_init %entry %s_next %loop_cont\n";
1101 o +=
"OpLoopMerge %loop_merge %loop_cont None\n";
1102 o +=
"OpBranch %loop_body\n\n";
1104 o +=
"%loop_body = OpLabel\n";
1105 o +=
"%active = OpULessThan %bool %lid %stride\n";
1106 o +=
"OpSelectionMerge %sel_merge None\n";
1107 o +=
"OpBranchConditional %active %do_op %sel_merge\n\n";
1109 o +=
"%do_op = OpLabel\n";
1110 o +=
"%pgsh_a = OpAccessChain %psh_f32 %shared_a %lid\n";
1111 o +=
"%a = OpLoad %f32 %pgsh_a\n";
1112 o +=
"%lid_b = OpIAdd %u32 %lid %stride\n";
1113 o +=
"%pgsh_b = OpAccessChain %psh_f32 %shared_a %lid_b\n";
1114 o +=
"%b = OpLoad %f32 %pgsh_b\n";
1115 o +=
"%pgidx_a = OpAccessChain %psh_u32 %shared_idx %lid\n";
1116 o +=
"%idx_a = OpLoad %u32 %pgidx_a\n";
1117 o +=
"%pgidx_b = OpAccessChain %psh_u32 %shared_idx %lid_b\n";
1118 o +=
"%idx_b = OpLoad %u32 %pgidx_b\n";
1119 o +=
"%b_wins = OpFOrdGreaterThan %bool %b %a\n";
1120 o +=
"%combined = OpSelect %f32 %b_wins %b %a\n";
1121 o +=
"%combined_idx = OpSelect %u32 %b_wins %idx_b %idx_a\n";
1122 o +=
"OpStore %pgsh_a %combined\n";
1123 o +=
"OpStore %pgidx_a %combined_idx\n";
1124 o +=
"OpBranch %sel_merge\n\n";
1126 o +=
"%sel_merge = OpLabel\n";
1127 o +=
"OpBranch %loop_cont\n\n";
1129 o +=
"%loop_cont = OpLabel\n";
1130 o +=
"OpControlBarrier %c2u %c2u %c264u\n";
1131 o +=
"%s_next = OpShiftRightLogical %u32 %stride %c1u\n";
1132 o +=
"%done = OpIEqual %bool %s_next %c0u\n";
1133 o +=
"OpBranchConditional %done %loop_merge %loop_hdr\n\n";
1135 o +=
"%loop_merge = OpLabel\n";
1136 o +=
"%is_zero = OpIEqual %bool %lid %c0u\n";
1137 o +=
"OpSelectionMerge %write_merge None\n";
1138 o +=
"OpBranchConditional %is_zero %do_write %write_merge\n\n";
1140 o +=
"%do_write = OpLabel\n";
1141 o +=
"%result_v = OpLoad %f32 %pgsh\n";
1142 o +=
"%result_i = OpLoad %u32 %pgidx\n";
1143 o +=
"%gep_out_v = OpAccessChain %pelem_" + b0.name
1144 +
" %buf_" + b0.name +
" %c0u %c0u\n";
1145 o +=
"OpStore %gep_out_v %result_v\n";
1146 o +=
"%gep_out_i = OpAccessChain %pelem_" + b1.name
1147 +
" %buf_" + b1.name +
" %c0u %c0u\n";
1148 o +=
"OpStore %gep_out_i %result_i\n";
1149 o +=
"OpBranch %write_merge\n\n";
1151 o +=
"%write_merge = OpLabel\n";
1153 o +=
"OpFunctionEnd\n";
1167 std::string emit_image_body(
const ShaderSpec& spec)
1170 std::vector<const BindingSlot*> img_inputs;
1177 img_inputs.push_back(&
b);
1182 o +=
"%main = OpFunction %void None %voidfn\n";
1183 o +=
"%entry = OpLabel\n";
1184 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1185 o +=
"%ix = OpCompositeExtract %u32 %gid3 0\n";
1186 o +=
"%iy = OpCompositeExtract %u32 %gid3 1\n";
1187 o +=
"%six = OpBitcast %i32 %ix\n";
1188 o +=
"%siy = OpBitcast %i32 %iy\n";
1189 o +=
"%coord = OpCompositeConstruct %v2i32 %six %siy\n\n";
1192 o +=
"%img_out_val = OpLoad %img2d_t %img_" + std::string(img_out->
name) +
"\n";
1193 for (
size_t ii = 0; ii < img_inputs.size(); ++ii)
1194 o +=
"%img_in_val" + std::to_string(ii) +
" = OpLoad %img2d_t %img_" + img_inputs[ii]->
name +
"\n";
1197 for (
size_t fi = 0; fi < spec.
pc_fields.size(); ++fi) {
1199 o +=
"%ppc_" + f.name +
" = OpAccessChain %ppc_f32 %pc %c"
1200 + std::to_string(fi) +
"u\n";
1201 o +=
"%pc_" + f.name +
" = OpLoad %f32 %ppc_" + f.name +
"\n";
1206 std::vector<const BindingSlot*> tex_inputs;
1209 tex_inputs.push_back(&
b);
1212 if (!tex_inputs.empty()) {
1213 const std::string pw = spec.
pc_fields.size() > 0
1216 const std::string ph = spec.
pc_fields.size() > 1
1219 o +=
"%fix = OpConvertUToF %f32 %ix\n";
1220 o +=
"%fiy = OpConvertUToF %f32 %iy\n";
1221 o +=
"%u = OpFDiv %f32 %fix " + pw +
"\n";
1222 o +=
"%v = OpFDiv %f32 %fiy " + ph +
"\n";
1223 o +=
"%uv = OpCompositeConstruct %v2f32 %u %v\n\n";
1225 for (
size_t ti = 0; ti < tex_inputs.size(); ++ti) {
1226 const std::string idx = std::to_string(img_inputs.size() + ti);
1227 o +=
"%simg_" + tex_inputs[ti]->name
1228 +
" = OpLoad %simgc_t %tex_" + tex_inputs[ti]->name +
"\n";
1229 o +=
"%raw_in" + idx +
" = OpImageSampleExplicitLod %v4f32 %simg_"
1230 + tex_inputs[ti]->name +
" %uv Lod %lod_zero\n";
1235 for (
size_t ii = 0; ii < img_inputs.size(); ++ii) {
1236 const std::string idx = std::to_string(ii);
1237 o +=
"%raw_in" + idx +
" = OpImageRead %v4f32 %img_in_val" + idx +
" %coord\n";
1239 if (img_inputs.empty()) {
1240 o +=
"%raw_in0 = OpCompositeConstruct %v4f32 %img_czero %img_czero"
1241 " %img_czero %img_czero\n";
1245 o +=
"%ch0_r = OpCompositeExtract %f32 %raw_in0 0\n";
1246 o +=
"%ch0_g = OpCompositeExtract %f32 %raw_in0 1\n";
1247 o +=
"%ch0_b = OpCompositeExtract %f32 %raw_in0 2\n";
1248 o +=
"%ch0_a = OpCompositeExtract %f32 %raw_in0 3\n";
1250 const bool has_second = img_inputs.size() > 1
1251 || (!img_inputs.empty() && !tex_inputs.empty())
1252 || tex_inputs.size() > 1;
1254 const std::string second_idx = img_inputs.size() > 1
1256 : (!tex_inputs.empty() ? std::to_string(img_inputs.size()) :
"");
1258 if (has_second && !second_idx.empty()) {
1259 o +=
"%ch1_r = OpCompositeExtract %f32 %raw_in" + second_idx +
" 0\n";
1260 o +=
"%ch1_g = OpCompositeExtract %f32 %raw_in" + second_idx +
" 1\n";
1261 o +=
"%ch1_b = OpCompositeExtract %f32 %raw_in" + second_idx +
" 2\n";
1262 o +=
"%ch1_a = OpCompositeExtract %f32 %raw_in" + second_idx +
" 3\n";
1266 const std::string p0 = spec.
pc_fields.empty()
1269 const std::string p1 = spec.
pc_fields.size() > 1
1274 const std::string&
wr =
"%pc_" + spec.
pc_fields[0].name;
1275 const std::string&
wg = spec.
pc_fields.size() > 1 ?
"%pc_" + spec.
pc_fields[1].name :
"%img_czero";
1276 const std::string&
wb = spec.
pc_fields.size() > 2 ?
"%pc_" + spec.
pc_fields[2].name :
"%img_czero";
1277 const std::string&
wa = spec.
pc_fields.size() > 3 ?
"%pc_" + spec.
pc_fields[3].name :
"%img_czero";
1278 o +=
"%dot_r = OpFMul %f32 %ch0_r " +
wr +
"\n";
1279 o +=
"%dot_g = OpFMul %f32 %ch0_g " +
wg +
"\n";
1280 o +=
"%dot_b = OpFMul %f32 %ch0_b " +
wb +
"\n";
1281 o +=
"%dot_a = OpFMul %f32 %ch0_a " +
wa +
"\n";
1282 o +=
"%dot_rg = OpFAdd %f32 %dot_r %dot_g\n";
1283 o +=
"%dot_rgb = OpFAdd %f32 %dot_rg %dot_b\n";
1284 o +=
"%dot_val = OpFAdd %f32 %dot_rgb %dot_a\n";
1285 o +=
"%out_vec = OpCompositeConstruct %v4f32 %dot_val %dot_val %dot_val %dot_val\n";
1287 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1289 o +=
"OpFunctionEnd\n";
1294 o +=
"%out_vec = OpCompositeConstruct %v4f32 %ch0_r %ch0_r %ch0_r %img_cone\n";
1296 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1298 o +=
"OpFunctionEnd\n";
1302 auto emit_channel_op = [&](
1303 const std::string& c0,
const std::string& c1,
1304 const std::string& suffix) {
1307 o +=
"%res_" + suffix +
" = OpFMul %f32 " + c0 +
" " + p0 +
"\n";
1310 o +=
"%mul_" + suffix +
" = OpFMul %f32 " + c0 +
" " + p0 +
"\n";
1311 o +=
"%res_" + suffix +
" = OpFAdd %f32 %mul_" + suffix +
" " + p1 +
"\n";
1314 o +=
"%res_" + suffix +
" = OpFAdd %f32 " + c0 +
" " + p0 +
"\n";
1317 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FClamp "
1318 + c0 +
" " + p0 +
" " + p1 +
"\n";
1321 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FAbs " + c0 +
"\n";
1324 o +=
"%res_" + suffix +
" = OpFNegate %f32 " + c0 +
"\n";
1327 o +=
"%res_" + suffix +
" = OpFAdd %f32 " + c0 +
" " + c1 +
"\n";
1330 o +=
"%res_" + suffix +
" = OpFMul %f32 " + c0 +
" " + c1 +
"\n";
1333 o +=
"%res_" + suffix +
" = OpExtInst %f32 %glsl FMix "
1334 + c0 +
" " + c1 +
" " + p0 +
"\n";
1337 o +=
"%res_" + suffix +
" = OpFSub %f32 " + c0 +
" " + c1 +
"\n";
1340 o +=
"%cmp_" + suffix +
" = OpFOrdGreaterThanEqual %bool " + c0 +
" " + p0 +
"\n";
1341 o +=
"%res_" + suffix +
" = OpSelect %f32 %cmp_" + suffix +
" %img_cone %img_czero\n";
1344 o +=
"%cmp_" + suffix +
" = OpFOrdGreaterThanEqual %bool " + c0 +
" " + p0 +
"\n";
1345 o +=
"%res_" + suffix +
" = OpSelect %f32 %cmp_" + suffix +
" " + p1 +
" " + c0 +
"\n";
1348 o +=
"%res_" + suffix +
" = OpCopyObject %f32 " + c0 +
"\n";
1353 const std::string zero =
"%img_czero";
1354 const std::string c1r = has_second ?
"%ch1_r" : zero;
1355 const std::string c1g = has_second ?
"%ch1_g" : zero;
1356 const std::string c1b = has_second ?
"%ch1_b" : zero;
1357 const std::string c1a = has_second ?
"%ch1_a" : zero;
1359 emit_channel_op(
"%ch0_r", c1r,
"r");
1360 emit_channel_op(
"%ch0_g", c1g,
"g");
1361 emit_channel_op(
"%ch0_b", c1b,
"b");
1362 emit_channel_op(
"%ch0_a", c1a,
"a");
1365 o +=
"%out_vec = OpCompositeConstruct %v4f32 %res_r %res_g %res_b %res_a\n";
1367 o +=
"OpImageWrite %img_out_val %coord %out_vec\n";
1370 o +=
"OpFunctionEnd\n";
1378 std::string emit_bitonic_body(
const ShaderSpec& spec)
1380 const auto& bkeys = spec.
bindings[0];
1381 const auto& bidx = spec.
bindings[1];
1383 const std::string ktype(ssbo_elem_spirv_type(bkeys.format));
1384 const std::string itype(ssbo_elem_spirv_type(bidx.format));
1387 o +=
"%main = OpFunction %void None %voidfn\n";
1388 o +=
"%entry = OpLabel\n";
1389 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1390 o +=
"%i = OpCompositeExtract %u32 %gid3 0\n\n";
1392 o +=
"%ppc_stage = OpAccessChain %ppc_u32 %pc %c0u\n";
1393 o +=
"%stage = OpLoad %u32 %ppc_stage\n";
1394 o +=
"%ppc_pass = OpAccessChain %ppc_u32 %pc %c1u\n";
1395 o +=
"%pass = OpLoad %u32 %ppc_pass\n";
1396 o +=
"%ppc_count = OpAccessChain %ppc_u32 %pc %c2u\n";
1397 o +=
"%count = OpLoad %u32 %ppc_count\n";
1398 o +=
"%ppc_desc = OpAccessChain %ppc_u32 %pc %c3u\n";
1399 o +=
"%descending = OpLoad %u32 %ppc_desc\n\n";
1401 o +=
"%c1u_shift = OpShiftLeftLogical %u32 %c1u %pass\n";
1402 o +=
"%partner = OpBitwiseXor %u32 %i %c1u_shift\n\n";
1404 o +=
"%partner_le_i = OpULessThanEqual %bool %partner %i\n";
1405 o +=
"%i_oob = OpUGreaterThanEqual %bool %i %count\n";
1406 o +=
"%p_oob = OpUGreaterThanEqual %bool %partner %count\n";
1407 o +=
"%oob_raw = OpLogicalOr %bool %i_oob %p_oob\n";
1408 o +=
"%skip = OpLogicalOr %bool %partner_le_i %oob_raw\n\n";
1410 o +=
"OpSelectionMerge %early_merge None\n";
1411 o +=
"OpBranchConditional %skip %early_ret %do_sort\n\n";
1413 o +=
"%do_sort = OpLabel\n";
1415 o +=
"%gep_ki = OpAccessChain %pelem_" + bkeys.name
1416 +
" %buf_" + bkeys.name +
" %c0u %i\n";
1417 o +=
"%key_i = OpLoad " + ktype +
" %gep_ki\n";
1418 o +=
"%gep_kp = OpAccessChain %pelem_" + bkeys.name
1419 +
" %buf_" + bkeys.name +
" %c0u %partner\n";
1420 o +=
"%key_p = OpLoad " + ktype +
" %gep_kp\n\n";
1422 o +=
"%gep_ii = OpAccessChain %pelem_" + bidx.name
1423 +
" %buf_" + bidx.name +
" %c0u %i\n";
1424 o +=
"%idx_i = OpLoad " + itype +
" %gep_ii\n";
1425 o +=
"%gep_ip = OpAccessChain %pelem_" + bidx.name
1426 +
" %buf_" + bidx.name +
" %c0u %partner\n";
1427 o +=
"%idx_p = OpLoad " + itype +
" %gep_ip\n\n";
1429 o +=
"%dir_shift = OpShiftRightLogical %u32 %i %stage\n";
1430 o +=
"%dir_bit = OpBitwiseAnd %u32 %dir_shift %c1u\n\n";
1432 o +=
"%gt = OpFOrdGreaterThan %bool %key_i %key_p\n";
1433 o +=
"%gt_u = OpSelect %u32 %gt %c1u %c0u\n";
1434 o +=
"%xor1 = OpBitwiseXor %u32 %gt_u %dir_bit\n";
1435 o +=
"%xor2 = OpBitwiseXor %u32 %xor1 %descending\n";
1436 o +=
"%do_swap = OpINotEqual %bool %xor2 %c0u\n\n";
1438 o +=
"%new_ki = OpSelect " + ktype +
" %do_swap %key_p %key_i\n";
1439 o +=
"%new_kp = OpSelect " + ktype +
" %do_swap %key_i %key_p\n";
1440 o +=
"%new_ii = OpSelect " + itype +
" %do_swap %idx_p %idx_i\n";
1441 o +=
"%new_ip = OpSelect " + itype +
" %do_swap %idx_i %idx_p\n\n";
1443 o +=
"OpStore %gep_ki %new_ki\n";
1444 o +=
"OpStore %gep_kp %new_kp\n";
1445 o +=
"OpStore %gep_ii %new_ii\n";
1446 o +=
"OpStore %gep_ip %new_ip\n";
1447 o +=
"OpBranch %early_merge\n\n";
1449 o +=
"%early_ret = OpLabel\n";
1450 o +=
"OpBranch %early_merge\n\n";
1452 o +=
"%early_merge = OpLabel\n";
1454 o +=
"OpFunctionEnd\n";
1463 std::string emit_convolve2d_body(
const ShaderSpec& spec)
1481 o +=
"%main = OpFunction %void None %voidfn\n";
1482 o +=
"%entry = OpLabel\n";
1484 o +=
"%gid3 = OpLoad %v3u32 %gid_var\n";
1485 o +=
"%ix = OpCompositeExtract %u32 %gid3 0\n";
1486 o +=
"%iy = OpCompositeExtract %u32 %gid3 1\n";
1488 o +=
"%ppc_radius = OpAccessChain %ppc_u32 %pc %c0u\n";
1489 o +=
"%pc_radius = OpLoad %u32 %ppc_radius\n";
1490 o +=
"%ppc_width = OpAccessChain %ppc_u32 %pc %c1u\n";
1491 o +=
"%pc_width = OpLoad %u32 %ppc_width\n";
1492 o +=
"%ppc_height = OpAccessChain %ppc_u32 %pc %c2u\n";
1493 o +=
"%pc_height = OpLoad %u32 %ppc_height\n";
1495 o +=
"%oob_x = OpUGreaterThanEqual %bool %ix %pc_width\n";
1496 o +=
"%oob_y = OpUGreaterThanEqual %bool %iy %pc_height\n";
1497 o +=
"%oob = OpLogicalOr %bool %oob_x %oob_y\n";
1498 o +=
"OpSelectionMerge %main_merge None\n";
1499 o +=
"OpBranchConditional %oob %main_merge %conv_start\n\n";
1501 o +=
"%conv_start = OpLabel\n";
1503 o +=
"%img_src_val = OpLoad %img2d_t %img_" + std::string(img_src->
name) +
"\n";
1504 o +=
"%img_out_val = OpLoad %img2d_t %img_" + std::string(img_out->
name) +
"\n";
1506 o +=
"%diam = OpIMul %u32 %pc_radius %c2u\n";
1507 o +=
"%diam1 = OpIAdd %u32 %diam %c1u\n";
1509 o +=
"%six = OpBitcast %i32 %ix\n";
1510 o +=
"%siy = OpBitcast %i32 %iy\n";
1511 o +=
"%srad = OpBitcast %i32 %pc_radius\n";
1512 o +=
"%sw = OpBitcast %i32 %pc_width\n";
1513 o +=
"%sh = OpBitcast %i32 %pc_height\n";
1514 o +=
"%sw_1 = OpISub %i32 %sw %ci1\n";
1515 o +=
"%sh_1 = OpISub %i32 %sh %ci1\n";
1517 o +=
"%czero4 = OpCompositeConstruct %v4f32 %img_czero %img_czero %img_czero %img_czero\n";
1519 o +=
"OpBranch %ky_hdr\n\n";
1521 o +=
"%ky_hdr = OpLabel\n";
1522 o +=
"%ky_u = OpPhi %u32 %c0u %conv_start %ky_next %ky_cont\n";
1523 o +=
"%acc_ky = OpPhi %v4f32 %czero4 %conv_start %acc_kx_done %ky_cont\n";
1524 o +=
"%ky_done = OpUGreaterThanEqual %bool %ky_u %diam1\n";
1525 o +=
"OpLoopMerge %ky_merge %ky_cont None\n";
1526 o +=
"OpBranchConditional %ky_done %ky_merge %kx_pre\n\n";
1528 o +=
"%kx_pre = OpLabel\n";
1529 o +=
"%ky_si = OpBitcast %i32 %ky_u\n";
1530 o +=
"%ky_off = OpISub %i32 %ky_si %srad\n";
1531 o +=
"%sy_raw = OpIAdd %i32 %siy %ky_off\n";
1532 o +=
"%sy_lo = OpExtInst %i32 %glsl SMax %sy_raw %ci0\n";
1533 o +=
"%sy = OpExtInst %i32 %glsl SMin %sy_lo %sh_1\n";
1534 o +=
"OpBranch %kx_hdr\n\n";
1536 o +=
"%kx_hdr = OpLabel\n";
1537 o +=
"%kx_u = OpPhi %u32 %c0u %kx_pre %kx_next %kx_cont\n";
1538 o +=
"%acc_kx = OpPhi %v4f32 %acc_ky %kx_pre %acc_new %kx_cont\n";
1539 o +=
"%kx_done = OpUGreaterThanEqual %bool %kx_u %diam1\n";
1540 o +=
"OpLoopMerge %kx_merge %kx_cont None\n";
1541 o +=
"OpBranchConditional %kx_done %kx_merge %kx_body\n\n";
1543 o +=
"%kx_body = OpLabel\n";
1544 o +=
"%kx_si = OpBitcast %i32 %kx_u\n";
1545 o +=
"%kx_off = OpISub %i32 %kx_si %srad\n";
1546 o +=
"%sx_raw = OpIAdd %i32 %six %kx_off\n";
1547 o +=
"%sx_lo = OpExtInst %i32 %glsl SMax %sx_raw %ci0\n";
1548 o +=
"%sx = OpExtInst %i32 %glsl SMin %sx_lo %sw_1\n";
1550 o +=
"%sc = OpCompositeConstruct %v2i32 %sx %sy\n";
1551 o +=
"%px = OpImageRead %v4f32 %img_src_val %sc\n";
1553 o +=
"%kidx_r = OpIMul %u32 %ky_u %diam1\n";
1554 o +=
"%kidx = OpIAdd %u32 %kidx_r %kx_u\n";
1555 o +=
"%k_gep = OpAccessChain %pelem_" + std::string(kern_ssbo->
name)
1556 +
" %buf_" + std::string(kern_ssbo->
name) +
" %c0u %kidx\n";
1557 o +=
"%kw = OpLoad %f32 %k_gep\n";
1559 o +=
"%kw4 = OpCompositeConstruct %v4f32 %kw %kw %kw %kw\n";
1560 o +=
"%prod = OpFMul %v4f32 %px %kw4\n";
1561 o +=
"%acc_new = OpFAdd %v4f32 %acc_kx %prod\n";
1562 o +=
"OpBranch %kx_cont\n\n";
1564 o +=
"%kx_cont = OpLabel\n";
1565 o +=
"%kx_next = OpIAdd %u32 %kx_u %c1u\n";
1566 o +=
"OpBranch %kx_hdr\n\n";
1568 o +=
"%kx_merge = OpLabel\n";
1569 o +=
"%acc_kx_done = OpPhi %v4f32 %acc_kx %kx_hdr\n";
1570 o +=
"OpBranch %ky_cont\n\n";
1572 o +=
"%ky_cont = OpLabel\n";
1573 o +=
"%ky_next = OpIAdd %u32 %ky_u %c1u\n";
1574 o +=
"OpBranch %ky_hdr\n\n";
1576 o +=
"%ky_merge = OpLabel\n";
1577 o +=
"%final_acc = OpPhi %v4f32 %acc_ky %ky_hdr\n";
1579 o +=
"%out_coord = OpCompositeConstruct %v2i32 %six %siy\n";
1580 o +=
"OpImageWrite %img_out_val %out_coord %final_acc\n";
1581 o +=
"OpBranch %main_merge\n\n";
1583 o +=
"%main_merge = OpLabel\n";
1585 o +=
"OpFunctionEnd\n";
1594 src += emit_header(spec);
1595 src += emit_decorations(spec);
1596 src += emit_types(spec);
1599 src += emit_convolve2d_body(spec);
1603 bool has_image =
false;
1612 src += emit_image_body(spec);
1616 switch (spec.
tmpl) {
1618 src += (spec.
op ==
KernelOp::MaxIndex) ? emit_max_index_body(spec) : emit_reduction_body(spec);
1621 src += emit_scan_body(spec);
1624 src += emit_bitonic_body(spec);
1627 src += emit_stencil_body(spec);
1632 src += emit_elementwise_body(spec);
1641 const auto& ks = *spec.
kernel;
1643 bool has_image =
false;
1650 o +=
"#version 460\n";
1651 o +=
"layout(local_size_x = " + std::to_string(ws[0])
1652 +
", local_size_y = " + std::to_string(ws[1])
1653 +
", local_size_z = " + std::to_string(ws[2]) +
") in;\n\n";
1660 o +=
"layout(set = 0, binding = " + std::to_string(
b.binding_index)
1661 +
", rgba32f) " + qual +
" uniform image2D " +
b.name +
";\n";
1665 o +=
"layout(set = 0, binding = " + std::to_string(
b.binding_index)
1666 +
") uniform sampler2D " +
b.name +
";\n";
1669 const auto t = std::string(glsl_type(
b.format));
1670 o +=
"layout(set = 0, binding = " + std::to_string(
b.binding_index)
1671 +
", std430) buffer Block_" +
b.name
1672 +
" { " + t +
" " +
b.name +
"[]; };\n";
1676 o +=
"\nlayout(push_constant) uniform PC {\n";
1678 o +=
" " + std::string(glsl_type(f.format)) +
" " + f.name +
";\n";
1683 o +=
"\n" + f.return_type +
" " + f.name +
"(" + f.params +
") {\n";
1688 o +=
"\nvoid main() {\n";
1689 o +=
" uint i = gl_GlobalInvocationID.x;\n";
1691 o +=
" ivec2 coord = ivec2(gl_GlobalInvocationID.xy);\n";
1694 const auto t = std::string(glsl_type(f.format));
1695 o +=
" " + t +
" " + f.name +
" = pc." + f.name +
";\n";