blob: fa4d4e3b27f0675f64202fce2ef67cbcba6b8eb0 [file] [edit]
#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);
}