Skip to content

Commit f803515

Browse files
kpetdneto0
andauthored
Add support for SPV_EXT_ocp_microscaling_types (KhronosGroup#6772)
Co-authored-by: Kevin Petit <kevin.petit@arm.com> Co-authored-by: Guillaume Trebuchet <guillaume.trebuchet@arm.com> --------- Signed-off-by: Kevin Petit <kevin.petit@arm.com> Signed-off-by: Guillaume Trebuchet <guillaume.trebuchet@arm.com> Co-authored-by: David Neto <dneto@google.com>
1 parent cbcd15a commit f803515

31 files changed

Lines changed: 2470 additions & 102 deletions

.github/workflows/bazel.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@ jobs:
3131
with:
3232
path: ~/.bazel/cache
3333
# Force a new cache version by updating the number in 'generationN'
34-
key: bazel-cache-${{ runner.os }}-generation2
34+
key: bazel-cache-${{ runner.os }}-generation3
3535
- name: Build All
3636
run: bazel --output_user_root=~/.bazel/cache build //...
3737
- name: Test All

DEPS

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ vars = {
1414

1515
're2_revision': '972a15cedd008d846f1a39b2e88ce48d7f166cbd',
1616

17-
'spirv_headers_revision': 'daa093dd29aab8cbb6562b808370562f56e399fb',
17+
'spirv_headers_revision': '575b6512579ebde466ed3dfc04e413439d14d95d',
1818

1919
'mimalloc_revision': 'fef6b0dd70f9d7fa0750b0d0b9fbb471203b94cd',
2020
}

include/spirv-tools/libspirv.h

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -398,6 +398,11 @@ typedef enum spv_fp_encoding_t {
398398
SPV_FP_ENCODING_BFLOAT16,
399399
SPV_FP_ENCODING_FLOAT8_E4M3,
400400
SPV_FP_ENCODING_FLOAT8_E5M2,
401+
SPV_FP_ENCODING_FLOAT6_E2M3,
402+
SPV_FP_ENCODING_FLOAT6_E3M2,
403+
SPV_FP_ENCODING_FLOAT4_E2M1,
404+
SPV_FP_ENCODING_FLOAT8_UNSIGNED_E8M0,
405+
SPV_FP_ENCODING_MXINT8,
401406
} spv_fp_encoding_t;
402407

403408
typedef enum spv_text_to_binary_options_t {

source/name_mapper.cpp

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -235,6 +235,27 @@ spv_result_t FriendlyNameMapper::ParseInstruction(
235235
SaveName(result_id, "fp8e5m2");
236236
break;
237237
}
238+
if (spv::FPEncoding(inst.words[3]) == spv::FPEncoding::Float6E2M3EXT) {
239+
SaveName(result_id, "fp6e2m3");
240+
break;
241+
}
242+
if (spv::FPEncoding(inst.words[3]) == spv::FPEncoding::Float6E3M2EXT) {
243+
SaveName(result_id, "fp6e3m2");
244+
break;
245+
}
246+
if (spv::FPEncoding(inst.words[3]) == spv::FPEncoding::Float4E2M1EXT) {
247+
SaveName(result_id, "fp4e2m1");
248+
break;
249+
}
250+
if (spv::FPEncoding(inst.words[3]) ==
251+
spv::FPEncoding::Float8UnsignedE8M0EXT) {
252+
SaveName(result_id, "fp8e8m0");
253+
break;
254+
}
255+
if (spv::FPEncoding(inst.words[3]) == spv::FPEncoding::MXInt8EXT) {
256+
SaveName(result_id, "mxint8");
257+
break;
258+
}
238259
}
239260
switch (bit_width) {
240261
case 16:

source/operand.cpp

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -647,6 +647,16 @@ spv_fp_encoding_t spvFPEncodingFromOperandFPEncoding(spv::FPEncoding encoding) {
647647
return SPV_FP_ENCODING_FLOAT8_E4M3;
648648
case spv::FPEncoding::Float8E5M2EXT:
649649
return SPV_FP_ENCODING_FLOAT8_E5M2;
650+
case spv::FPEncoding::Float6E2M3EXT:
651+
return SPV_FP_ENCODING_FLOAT6_E2M3;
652+
case spv::FPEncoding::Float6E3M2EXT:
653+
return SPV_FP_ENCODING_FLOAT6_E3M2;
654+
case spv::FPEncoding::Float4E2M1EXT:
655+
return SPV_FP_ENCODING_FLOAT4_E2M1;
656+
case spv::FPEncoding::Float8UnsignedE8M0EXT:
657+
return SPV_FP_ENCODING_FLOAT8_UNSIGNED_E8M0;
658+
case spv::FPEncoding::MXInt8EXT:
659+
return SPV_FP_ENCODING_MXINT8;
650660
case spv::FPEncoding::Max:
651661
break;
652662
}

source/opt/type_manager.cpp

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -248,11 +248,19 @@ uint32_t TypeManager::GetTypeInstruction(const Type* type) {
248248
{(type->AsInteger()->IsSigned() ? 1u : 0u)}}});
249249
break;
250250
case Type::kFloat:
251-
// TODO: Handle FP encoding enums once actually used.
252-
typeInst = MakeUnique<Instruction>(
253-
context(), spv::Op::OpTypeFloat, 0, id,
254-
std::initializer_list<Operand>{
255-
{SPV_OPERAND_TYPE_LITERAL_INTEGER, {type->AsFloat()->width()}}});
251+
if (type->AsFloat()->encoding() == spv::FPEncoding::Max) {
252+
typeInst = MakeUnique<Instruction>(
253+
context(), spv::Op::OpTypeFloat, 0, id,
254+
std::initializer_list<Operand>{{SPV_OPERAND_TYPE_LITERAL_INTEGER,
255+
{type->AsFloat()->width()}}});
256+
} else {
257+
typeInst = MakeUnique<Instruction>(
258+
context(), spv::Op::OpTypeFloat, 0, id,
259+
std::initializer_list<Operand>{
260+
{SPV_OPERAND_TYPE_LITERAL_INTEGER, {type->AsFloat()->width()}},
261+
{SPV_OPERAND_TYPE_FPENCODING,
262+
{static_cast<uint32_t>(type->AsFloat()->encoding())}}});
263+
}
256264
break;
257265
case Type::kVector: {
258266
uint32_t subtype = GetTypeInstruction(type->AsVector()->element_type());

source/opt/types.cpp

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -400,6 +400,18 @@ bool Float::IsSameImpl(const Type* that, IsSameCache*) const {
400400
std::string Float::str() const {
401401
std::ostringstream oss;
402402
switch (encoding_) {
403+
case spv::FPEncoding::Float4E2M1EXT:
404+
assert(width_ == 4);
405+
oss << "fp4e2m1";
406+
break;
407+
case spv::FPEncoding::Float6E2M3EXT:
408+
assert(width_ == 6);
409+
oss << "fp6e2m3";
410+
break;
411+
case spv::FPEncoding::Float6E3M2EXT:
412+
assert(width_ == 6);
413+
oss << "fp6e3m2";
414+
break;
403415
case spv::FPEncoding::BFloat16KHR:
404416
assert(width_ == 16);
405417
oss << "bfloat16";
@@ -412,6 +424,14 @@ std::string Float::str() const {
412424
assert(width_ == 8);
413425
oss << "fp8e5m2";
414426
break;
427+
case spv::FPEncoding::Float8UnsignedE8M0EXT:
428+
assert(width_ == 8);
429+
oss << "fp8e8m0";
430+
break;
431+
case spv::FPEncoding::MXInt8EXT:
432+
assert(width_ == 8);
433+
oss << "mxint8";
434+
break;
415435
default:
416436
oss << "float" << width_;
417437
break;

source/parsed_operand.cpp

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,18 @@ void EmitNumericLiteral(std::ostream* out, const spv_parsed_instruction_t& inst,
5151
case SPV_FP_ENCODING_IEEE754_BINARY32:
5252
*out << spvtools::utils::FloatProxy<float>(word);
5353
break;
54+
case SPV_FP_ENCODING_FLOAT4_E2M1:
55+
*out << spvtools::utils::FloatProxy<spvtools::utils::Float4_E2M1>(
56+
uint8_t(word & 0xF));
57+
break;
58+
case SPV_FP_ENCODING_FLOAT6_E2M3:
59+
*out << spvtools::utils::FloatProxy<spvtools::utils::Float6_E2M3>(
60+
uint8_t(word & 0x3F));
61+
break;
62+
case SPV_FP_ENCODING_FLOAT6_E3M2:
63+
*out << spvtools::utils::FloatProxy<spvtools::utils::Float6_E3M2>(
64+
uint8_t(word & 0x3F));
65+
break;
5466
case SPV_FP_ENCODING_FLOAT8_E4M3:
5567
*out << spvtools::utils::FloatProxy<spvtools::utils::Float8_E4M3>(
5668
uint8_t(word & 0xFF));
@@ -59,6 +71,14 @@ void EmitNumericLiteral(std::ostream* out, const spv_parsed_instruction_t& inst,
5971
*out << spvtools::utils::FloatProxy<spvtools::utils::Float8_E5M2>(
6072
uint8_t(word & 0xFF));
6173
break;
74+
case SPV_FP_ENCODING_FLOAT8_UNSIGNED_E8M0:
75+
*out << spvtools::utils::FloatProxy<spvtools::utils::Float8_E8M0>(
76+
uint8_t(word & 0xFF));
77+
break;
78+
case SPV_FP_ENCODING_MXINT8:
79+
*out << spvtools::utils::HexFixedPoint<spvtools::utils::MXInt8>(
80+
uint8_t(word & 0xFF));
81+
break;
6282
case SPV_FP_ENCODING_BFLOAT16:
6383
*out << spvtools::utils::FloatProxy<spvtools::utils::BFloat16>(
6484
uint16_t(word & 0xFFFF));

source/text_handler.cpp

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -168,8 +168,17 @@ bool validBitWidthForFPEncoding(spv_fp_encoding_t enc, uint32_t width,
168168
case SPV_FP_ENCODING_IEEE754_BINARY64:
169169
*expected = 64;
170170
break;
171+
case SPV_FP_ENCODING_FLOAT4_E2M1:
172+
*expected = 4;
173+
break;
174+
case SPV_FP_ENCODING_FLOAT6_E2M3:
175+
case SPV_FP_ENCODING_FLOAT6_E3M2:
176+
*expected = 6;
177+
break;
171178
case SPV_FP_ENCODING_FLOAT8_E5M2:
172179
case SPV_FP_ENCODING_FLOAT8_E4M3:
180+
case SPV_FP_ENCODING_FLOAT8_UNSIGNED_E8M0:
181+
case SPV_FP_ENCODING_MXINT8:
173182
*expected = 8;
174183
break;
175184
default:

0 commit comments

Comments
 (0)