| #version 450 core | |
| #extension GL_KHR_memory_scope_semantics : enable | |
| #extension GL_KHR_cooperative_matrix : enable | |
| #extension GL_EXT_shader_explicit_arithmetic_types : enable | |
| #extension GL_EXT_long_vector : enable | |
| #extension GL_NV_cooperative_matrix2 : enable | |
| #extension GL_NV_cooperative_matrix_decode_vector : enable | |
| #extension GL_EXT_buffer_reference : enable | |
| layout (local_size_x = 64, local_size_y = 1, local_size_z = 1) in; | |
| buffer BufType { | |
| float16_t x[]; | |
| } Buf; | |
| buffer BufType8 { | |
| uint8_t x[]; | |
| } Buf8; | |
| layout(buffer_reference, std430, buffer_reference_align = 4) buffer dwordBuf { | |
| uint32_t d; | |
| }; | |
| float16_t decodeF16Scalar(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return float16_t(b.d); | |
| } | |
| f16vec2 decodeF16x2(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return unpackFloat2x16(b.d); | |
| } | |
| f16vec4 decodeF16x4(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return f16vec4(unpackFloat2x16(b.d), unpackFloat2x16(b.d)); | |
| } | |
| uint8_t decodeU8Scalar(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return uint8_t(b.d); | |
| } | |
| u8vec4 decodeU8x4(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return u8vec4(uint8_t(b.d), uint8_t(b.d >> 8), uint8_t(b.d >> 16), uint8_t(b.d >> 24)); | |
| } | |
| vector<uint8_t, 8> decodeU8x8(const in dwordBuf b, const in uint32_t blockCoords[2], const in uint32_t coordInBlock[2]) | |
| { | |
| return vector<uint8_t, 8>(uint8_t(b.d), uint8_t(b.d >> 8), | |
| uint8_t(b.d >> 16), uint8_t(b.d >> 24), | |
| uint8_t(b.d), uint8_t(b.d >> 8), | |
| uint8_t(b.d >> 16), uint8_t(b.d >> 24)); | |
| } | |
| void main() | |
| { | |
| coopmat<float16_t, gl_ScopeWorkgroup, 64, 32, gl_MatrixUseA> A; | |
| coopmat<uint8_t, gl_ScopeWorkgroup, 64, 32, gl_MatrixUseA> Au8; | |
| tensorLayoutNV<2> t = createTensorLayoutNV(2); | |
| t = setTensorLayoutBlockSizeNV(t, 4, 8); | |
| t = setTensorLayoutDimensionNV(t, 256, 512); | |
| // Scalar-only decode (the SPV_NV_cooperative_matrix2 baseline). No | |
| // capability or extension from this extension is required when the | |
| // vector form is not used. | |
| coopMatLoadTensorNV(A, Buf.x, 0, t, decodeF16Scalar); | |
| coopMatLoadTensorNV(Au8, Buf8.x, 0, t, decodeU8Scalar); | |
| // Scalar + vector decode pair: the implementation may invoke either | |
| // function per call site. V == 2 fp16. | |
| coopMatLoadTensorNV(A, Buf.x, 0, t, decodeF16Scalar, decodeF16x2); | |
| // V == 4 fp16. | |
| coopMatLoadTensorNV(A, Buf.x, 0, t, decodeF16Scalar, decodeF16x4); | |
| // V == 4 u8. | |
| coopMatLoadTensorNV(Au8, Buf8.x, 0, t, decodeU8Scalar, decodeU8x4); | |
| // V == 8 u8. | |
| coopMatLoadTensorNV(Au8, Buf8.x, 0, t, decodeU8Scalar, decodeU8x8); | |
| } |