emit_spirv_context_get_set.cpp 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418
  1. // Copyright 2021 yuzu Emulator Project
  2. // Licensed under GPLv2 or any later version
  3. // Refer to the license.txt file included.
  4. #include <tuple>
  5. #include <utility>
  6. #include "shader_recompiler/backend/spirv/emit_spirv.h"
  7. namespace Shader::Backend::SPIRV {
  8. namespace {
  9. struct AttrInfo {
  10. Id pointer;
  11. Id id;
  12. bool needs_cast;
  13. };
  14. std::optional<AttrInfo> AttrTypes(EmitContext& ctx, u32 index) {
  15. const AttributeType type{ctx.profile.generic_input_types.at(index)};
  16. switch (type) {
  17. case AttributeType::Float:
  18. return AttrInfo{ctx.input_f32, ctx.F32[1], false};
  19. case AttributeType::UnsignedInt:
  20. return AttrInfo{ctx.input_u32, ctx.U32[1], true};
  21. case AttributeType::SignedInt:
  22. return AttrInfo{ctx.input_s32, ctx.TypeInt(32, true), true};
  23. case AttributeType::Disabled:
  24. return std::nullopt;
  25. }
  26. throw InvalidArgument("Invalid attribute type {}", type);
  27. }
  28. template <typename... Args>
  29. Id AttrPointer(EmitContext& ctx, Id pointer_type, Id vertex, Id base, Args&&... args) {
  30. switch (ctx.stage) {
  31. case Stage::TessellationControl:
  32. case Stage::TessellationEval:
  33. case Stage::Geometry:
  34. return ctx.OpAccessChain(pointer_type, base, vertex, std::forward<Args>(args)...);
  35. default:
  36. return ctx.OpAccessChain(pointer_type, base, std::forward<Args>(args)...);
  37. }
  38. }
  39. template <typename... Args>
  40. Id OutputAccessChain(EmitContext& ctx, Id result_type, Id base, Args&&... args) {
  41. if (ctx.stage == Stage::TessellationControl) {
  42. const Id invocation_id{ctx.OpLoad(ctx.U32[1], ctx.invocation_id)};
  43. return ctx.OpAccessChain(result_type, base, invocation_id, std::forward<Args>(args)...);
  44. } else {
  45. return ctx.OpAccessChain(result_type, base, std::forward<Args>(args)...);
  46. }
  47. }
  48. struct OutAttr {
  49. OutAttr(Id pointer_) : pointer{pointer_} {}
  50. OutAttr(Id pointer_, Id type_) : pointer{pointer_}, type{type_} {}
  51. Id pointer{};
  52. Id type{};
  53. };
  54. std::optional<OutAttr> OutputAttrPointer(EmitContext& ctx, IR::Attribute attr) {
  55. if (IR::IsGeneric(attr)) {
  56. const u32 index{IR::GenericAttributeIndex(attr)};
  57. const u32 element{IR::GenericAttributeElement(attr)};
  58. const GenericElementInfo& info{ctx.output_generics.at(index).at(element)};
  59. if (info.num_components == 1) {
  60. return info.id;
  61. } else {
  62. const u32 index_element{element - info.first_element};
  63. const Id index_id{ctx.Const(index_element)};
  64. return OutputAccessChain(ctx, ctx.output_f32, info.id, index_id);
  65. }
  66. }
  67. switch (attr) {
  68. case IR::Attribute::PointSize:
  69. return ctx.output_point_size;
  70. case IR::Attribute::PositionX:
  71. case IR::Attribute::PositionY:
  72. case IR::Attribute::PositionZ:
  73. case IR::Attribute::PositionW: {
  74. const u32 element{static_cast<u32>(attr) % 4};
  75. const Id element_id{ctx.Const(element)};
  76. return OutputAccessChain(ctx, ctx.output_f32, ctx.output_position, element_id);
  77. }
  78. case IR::Attribute::ClipDistance0:
  79. case IR::Attribute::ClipDistance1:
  80. case IR::Attribute::ClipDistance2:
  81. case IR::Attribute::ClipDistance3:
  82. case IR::Attribute::ClipDistance4:
  83. case IR::Attribute::ClipDistance5:
  84. case IR::Attribute::ClipDistance6:
  85. case IR::Attribute::ClipDistance7: {
  86. const u32 base{static_cast<u32>(IR::Attribute::ClipDistance0)};
  87. const u32 index{static_cast<u32>(attr) - base};
  88. const Id clip_num{ctx.Const(index)};
  89. return OutputAccessChain(ctx, ctx.output_f32, ctx.clip_distances, clip_num);
  90. }
  91. case IR::Attribute::Layer:
  92. if (ctx.profile.support_viewport_index_layer_non_geometry ||
  93. ctx.stage == Shader::Stage::Geometry) {
  94. return OutAttr{ctx.layer, ctx.U32[1]};
  95. }
  96. return std::nullopt;
  97. case IR::Attribute::ViewportIndex:
  98. if (ctx.profile.support_viewport_index_layer_non_geometry ||
  99. ctx.stage == Shader::Stage::Geometry) {
  100. return OutAttr{ctx.viewport_index, ctx.U32[1]};
  101. }
  102. return std::nullopt;
  103. case IR::Attribute::ViewportMask:
  104. if (!ctx.profile.support_viewport_mask) {
  105. return std::nullopt;
  106. }
  107. return OutAttr{ctx.OpAccessChain(ctx.output_u32, ctx.viewport_mask, ctx.u32_zero_value),
  108. ctx.U32[1]};
  109. default:
  110. throw NotImplementedException("Read attribute {}", attr);
  111. }
  112. }
  113. Id GetCbuf(EmitContext& ctx, Id result_type, Id UniformDefinitions::*member_ptr, u32 element_size,
  114. const IR::Value& binding, const IR::Value& offset) {
  115. if (!binding.IsImmediate()) {
  116. throw NotImplementedException("Constant buffer indexing");
  117. }
  118. const Id cbuf{ctx.cbufs[binding.U32()].*member_ptr};
  119. const Id uniform_type{ctx.uniform_types.*member_ptr};
  120. if (!offset.IsImmediate()) {
  121. Id index{ctx.Def(offset)};
  122. if (element_size > 1) {
  123. const u32 log2_element_size{static_cast<u32>(std::countr_zero(element_size))};
  124. const Id shift{ctx.Const(log2_element_size)};
  125. index = ctx.OpShiftRightArithmetic(ctx.U32[1], ctx.Def(offset), shift);
  126. }
  127. const Id access_chain{ctx.OpAccessChain(uniform_type, cbuf, ctx.u32_zero_value, index)};
  128. return ctx.OpLoad(result_type, access_chain);
  129. }
  130. if (offset.U32() % element_size != 0) {
  131. throw NotImplementedException("Unaligned immediate constant buffer load");
  132. }
  133. const Id imm_offset{ctx.Const(offset.U32() / element_size)};
  134. const Id access_chain{ctx.OpAccessChain(uniform_type, cbuf, ctx.u32_zero_value, imm_offset)};
  135. return ctx.OpLoad(result_type, access_chain);
  136. }
  137. } // Anonymous namespace
  138. void EmitGetRegister(EmitContext&) {
  139. throw NotImplementedException("SPIR-V Instruction");
  140. }
  141. void EmitSetRegister(EmitContext&) {
  142. throw NotImplementedException("SPIR-V Instruction");
  143. }
  144. void EmitGetPred(EmitContext&) {
  145. throw NotImplementedException("SPIR-V Instruction");
  146. }
  147. void EmitSetPred(EmitContext&) {
  148. throw NotImplementedException("SPIR-V Instruction");
  149. }
  150. void EmitSetGotoVariable(EmitContext&) {
  151. throw NotImplementedException("SPIR-V Instruction");
  152. }
  153. void EmitGetGotoVariable(EmitContext&) {
  154. throw NotImplementedException("SPIR-V Instruction");
  155. }
  156. void EmitSetIndirectBranchVariable(EmitContext&) {
  157. throw NotImplementedException("SPIR-V Instruction");
  158. }
  159. void EmitGetIndirectBranchVariable(EmitContext&) {
  160. throw NotImplementedException("SPIR-V Instruction");
  161. }
  162. Id EmitGetCbufU8(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  163. const Id load{GetCbuf(ctx, ctx.U8, &UniformDefinitions::U8, sizeof(u8), binding, offset)};
  164. return ctx.OpUConvert(ctx.U32[1], load);
  165. }
  166. Id EmitGetCbufS8(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  167. const Id load{GetCbuf(ctx, ctx.S8, &UniformDefinitions::S8, sizeof(s8), binding, offset)};
  168. return ctx.OpSConvert(ctx.U32[1], load);
  169. }
  170. Id EmitGetCbufU16(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  171. const Id load{GetCbuf(ctx, ctx.U16, &UniformDefinitions::U16, sizeof(u16), binding, offset)};
  172. return ctx.OpUConvert(ctx.U32[1], load);
  173. }
  174. Id EmitGetCbufS16(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  175. const Id load{GetCbuf(ctx, ctx.S16, &UniformDefinitions::S16, sizeof(s16), binding, offset)};
  176. return ctx.OpSConvert(ctx.U32[1], load);
  177. }
  178. Id EmitGetCbufU32(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  179. return GetCbuf(ctx, ctx.U32[1], &UniformDefinitions::U32, sizeof(u32), binding, offset);
  180. }
  181. Id EmitGetCbufF32(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  182. return GetCbuf(ctx, ctx.F32[1], &UniformDefinitions::F32, sizeof(f32), binding, offset);
  183. }
  184. Id EmitGetCbufU32x2(EmitContext& ctx, const IR::Value& binding, const IR::Value& offset) {
  185. return GetCbuf(ctx, ctx.U32[2], &UniformDefinitions::U32x2, sizeof(u32[2]), binding, offset);
  186. }
  187. Id EmitGetAttribute(EmitContext& ctx, IR::Attribute attr, Id vertex) {
  188. const u32 element{static_cast<u32>(attr) % 4};
  189. const auto element_id{[&] { return ctx.Const(element); }};
  190. if (IR::IsGeneric(attr)) {
  191. const u32 index{IR::GenericAttributeIndex(attr)};
  192. const std::optional<AttrInfo> type{AttrTypes(ctx, index)};
  193. if (!type) {
  194. // Attribute is disabled
  195. return ctx.Const(0.0f);
  196. }
  197. const Id generic_id{ctx.input_generics.at(index)};
  198. const Id pointer{AttrPointer(ctx, type->pointer, vertex, generic_id, element_id())};
  199. const Id value{ctx.OpLoad(type->id, pointer)};
  200. return type->needs_cast ? ctx.OpBitcast(ctx.F32[1], value) : value;
  201. }
  202. switch (attr) {
  203. case IR::Attribute::PrimitiveId:
  204. return ctx.OpBitcast(ctx.F32[1], ctx.OpLoad(ctx.U32[1], ctx.primitive_id));
  205. case IR::Attribute::PositionX:
  206. case IR::Attribute::PositionY:
  207. case IR::Attribute::PositionZ:
  208. case IR::Attribute::PositionW:
  209. return ctx.OpLoad(
  210. ctx.F32[1], AttrPointer(ctx, ctx.input_f32, vertex, ctx.input_position, element_id()));
  211. case IR::Attribute::InstanceId:
  212. if (ctx.profile.support_vertex_instance_id) {
  213. return ctx.OpBitcast(ctx.F32[1], ctx.OpLoad(ctx.U32[1], ctx.instance_id));
  214. } else {
  215. const Id index{ctx.OpLoad(ctx.U32[1], ctx.instance_index)};
  216. const Id base{ctx.OpLoad(ctx.U32[1], ctx.base_instance)};
  217. return ctx.OpBitcast(ctx.F32[1], ctx.OpISub(ctx.U32[1], index, base));
  218. }
  219. case IR::Attribute::VertexId:
  220. if (ctx.profile.support_vertex_instance_id) {
  221. return ctx.OpBitcast(ctx.F32[1], ctx.OpLoad(ctx.U32[1], ctx.vertex_id));
  222. } else {
  223. const Id index{ctx.OpLoad(ctx.U32[1], ctx.vertex_index)};
  224. const Id base{ctx.OpLoad(ctx.U32[1], ctx.base_vertex)};
  225. return ctx.OpBitcast(ctx.F32[1], ctx.OpISub(ctx.U32[1], index, base));
  226. }
  227. case IR::Attribute::FrontFace:
  228. return ctx.OpSelect(ctx.U32[1], ctx.OpLoad(ctx.U1, ctx.front_face),
  229. ctx.Const(std::numeric_limits<u32>::max()), ctx.u32_zero_value);
  230. case IR::Attribute::PointSpriteS:
  231. return ctx.OpLoad(ctx.F32[1],
  232. ctx.OpAccessChain(ctx.input_f32, ctx.point_coord, ctx.u32_zero_value));
  233. case IR::Attribute::PointSpriteT:
  234. return ctx.OpLoad(ctx.F32[1],
  235. ctx.OpAccessChain(ctx.input_f32, ctx.point_coord, ctx.Const(1U)));
  236. case IR::Attribute::TessellationEvaluationPointU:
  237. return ctx.OpLoad(ctx.F32[1],
  238. ctx.OpAccessChain(ctx.input_f32, ctx.tess_coord, ctx.u32_zero_value));
  239. case IR::Attribute::TessellationEvaluationPointV:
  240. return ctx.OpLoad(ctx.F32[1],
  241. ctx.OpAccessChain(ctx.input_f32, ctx.tess_coord, ctx.Const(1U)));
  242. default:
  243. throw NotImplementedException("Read attribute {}", attr);
  244. }
  245. }
  246. void EmitSetAttribute(EmitContext& ctx, IR::Attribute attr, Id value, [[maybe_unused]] Id vertex) {
  247. const std::optional<OutAttr> output{OutputAttrPointer(ctx, attr)};
  248. if (!output) {
  249. return;
  250. }
  251. if (Sirit::ValidId(output->type)) {
  252. value = ctx.OpBitcast(output->type, value);
  253. }
  254. ctx.OpStore(output->pointer, value);
  255. }
  256. Id EmitGetAttributeIndexed(EmitContext& ctx, Id offset, Id vertex) {
  257. switch (ctx.stage) {
  258. case Stage::TessellationControl:
  259. case Stage::TessellationEval:
  260. case Stage::Geometry:
  261. return ctx.OpFunctionCall(ctx.F32[1], ctx.indexed_load_func, offset, vertex);
  262. default:
  263. return ctx.OpFunctionCall(ctx.F32[1], ctx.indexed_load_func, offset);
  264. }
  265. }
  266. void EmitSetAttributeIndexed(EmitContext& ctx, Id offset, Id value, [[maybe_unused]] Id vertex) {
  267. ctx.OpFunctionCall(ctx.void_id, ctx.indexed_store_func, offset, value);
  268. }
  269. Id EmitGetPatch(EmitContext& ctx, IR::Patch patch) {
  270. if (!IR::IsGeneric(patch)) {
  271. throw NotImplementedException("Non-generic patch load");
  272. }
  273. const u32 index{IR::GenericPatchIndex(patch)};
  274. const Id element{ctx.Const(IR::GenericPatchElement(patch))};
  275. const Id pointer{ctx.OpAccessChain(ctx.input_f32, ctx.patches.at(index), element)};
  276. return ctx.OpLoad(ctx.F32[1], pointer);
  277. }
  278. void EmitSetPatch(EmitContext& ctx, IR::Patch patch, Id value) {
  279. const Id pointer{[&] {
  280. if (IR::IsGeneric(patch)) {
  281. const u32 index{IR::GenericPatchIndex(patch)};
  282. const Id element{ctx.Const(IR::GenericPatchElement(patch))};
  283. return ctx.OpAccessChain(ctx.output_f32, ctx.patches.at(index), element);
  284. }
  285. switch (patch) {
  286. case IR::Patch::TessellationLodLeft:
  287. case IR::Patch::TessellationLodRight:
  288. case IR::Patch::TessellationLodTop:
  289. case IR::Patch::TessellationLodBottom: {
  290. const u32 index{static_cast<u32>(patch) - u32(IR::Patch::TessellationLodLeft)};
  291. const Id index_id{ctx.Const(index)};
  292. return ctx.OpAccessChain(ctx.output_f32, ctx.output_tess_level_outer, index_id);
  293. }
  294. case IR::Patch::TessellationLodInteriorU:
  295. return ctx.OpAccessChain(ctx.output_f32, ctx.output_tess_level_inner,
  296. ctx.u32_zero_value);
  297. case IR::Patch::TessellationLodInteriorV:
  298. return ctx.OpAccessChain(ctx.output_f32, ctx.output_tess_level_inner, ctx.Const(1u));
  299. default:
  300. throw NotImplementedException("Patch {}", patch);
  301. }
  302. }()};
  303. ctx.OpStore(pointer, value);
  304. }
  305. void EmitSetFragColor(EmitContext& ctx, u32 index, u32 component, Id value) {
  306. const Id component_id{ctx.Const(component)};
  307. const Id pointer{ctx.OpAccessChain(ctx.output_f32, ctx.frag_color.at(index), component_id)};
  308. ctx.OpStore(pointer, value);
  309. }
  310. void EmitSetSampleMask(EmitContext& ctx, Id value) {
  311. ctx.OpStore(ctx.sample_mask, value);
  312. }
  313. void EmitSetFragDepth(EmitContext& ctx, Id value) {
  314. ctx.OpStore(ctx.frag_depth, value);
  315. }
  316. void EmitGetZFlag(EmitContext&) {
  317. throw NotImplementedException("SPIR-V Instruction");
  318. }
  319. void EmitGetSFlag(EmitContext&) {
  320. throw NotImplementedException("SPIR-V Instruction");
  321. }
  322. void EmitGetCFlag(EmitContext&) {
  323. throw NotImplementedException("SPIR-V Instruction");
  324. }
  325. void EmitGetOFlag(EmitContext&) {
  326. throw NotImplementedException("SPIR-V Instruction");
  327. }
  328. void EmitSetZFlag(EmitContext&) {
  329. throw NotImplementedException("SPIR-V Instruction");
  330. }
  331. void EmitSetSFlag(EmitContext&) {
  332. throw NotImplementedException("SPIR-V Instruction");
  333. }
  334. void EmitSetCFlag(EmitContext&) {
  335. throw NotImplementedException("SPIR-V Instruction");
  336. }
  337. void EmitSetOFlag(EmitContext&) {
  338. throw NotImplementedException("SPIR-V Instruction");
  339. }
  340. Id EmitWorkgroupId(EmitContext& ctx) {
  341. return ctx.OpLoad(ctx.U32[3], ctx.workgroup_id);
  342. }
  343. Id EmitLocalInvocationId(EmitContext& ctx) {
  344. return ctx.OpLoad(ctx.U32[3], ctx.local_invocation_id);
  345. }
  346. Id EmitInvocationId(EmitContext& ctx) {
  347. return ctx.OpLoad(ctx.U32[1], ctx.invocation_id);
  348. }
  349. Id EmitSampleId(EmitContext& ctx) {
  350. return ctx.OpLoad(ctx.U32[1], ctx.sample_id);
  351. }
  352. Id EmitIsHelperInvocation(EmitContext& ctx) {
  353. return ctx.OpLoad(ctx.U1, ctx.is_helper_invocation);
  354. }
  355. Id EmitYDirection(EmitContext& ctx) {
  356. return ctx.Const(ctx.profile.y_negate ? -1.0f : 1.0f);
  357. }
  358. Id EmitLoadLocal(EmitContext& ctx, Id word_offset) {
  359. const Id pointer{ctx.OpAccessChain(ctx.private_u32, ctx.local_memory, word_offset)};
  360. return ctx.OpLoad(ctx.U32[1], pointer);
  361. }
  362. void EmitWriteLocal(EmitContext& ctx, Id word_offset, Id value) {
  363. const Id pointer{ctx.OpAccessChain(ctx.private_u32, ctx.local_memory, word_offset)};
  364. ctx.OpStore(pointer, value);
  365. }
  366. } // namespace Shader::Backend::SPIRV