Browse Source

Implement Zvfbqdot8f

pull/2092/head
Andrew Waterman 10 months ago
parent
commit
e3dc14a878
  1. 2
      disasm/isa_parser.cc
  2. 17
      riscv/insns/vfqbdot_alt_vv.h
  3. 17
      riscv/insns/vfqbdot_vv.h
  4. 1
      riscv/isa_parser.h
  5. 2
      riscv/riscv.mk.in
  6. 2
      riscv/vector_unit.cc
  7. 15
      riscv/zvbdot.h

2
disasm/isa_parser.cc

@ -330,6 +330,8 @@ isa_parser_t::isa_parser_t(const char* str, const char *priv)
extension_table[EXT_ZVQBDOT8I] = true;
} else if (ext_str == "zvqbdot16i") {
extension_table[EXT_ZVQBDOT16I] = true;
} else if (ext_str == "zvfqbdot8f") {
extension_table[EXT_ZVFQBDOT8F] = true;
} else if (ext_str == "zvfwbdot16bf") {
extension_table[EXT_ZVFWBDOT16BF] = true;
} else if (ext_str == "zvfbdot32f") {

17
riscv/insns/vfqbdot_alt_vv.h

@ -0,0 +1,17 @@
VI_VFP_BASE;
ZVBDOT_INIT(4);
#define COMMA ,
switch (P.VU.vsew) {
case 8: {
require_extension(EXT_ZVFQBDOT8F);
if (P.VU.altfmt) {
ZVBDOT_LOOP(uint8_t, uint8_t, float32_t, zvfqbdot8f_dot_acc<ofp8_e5m2 COMMA ofp8_e5m2>);
} else {
ZVBDOT_LOOP(uint8_t, uint8_t, float32_t, zvfqbdot8f_dot_acc<ofp8_e4m3 COMMA ofp8_e5m2>);
}
break;
}
default: require(false);
}

17
riscv/insns/vfqbdot_vv.h

@ -0,0 +1,17 @@
VI_VFP_BASE;
ZVBDOT_INIT(4);
#define COMMA ,
switch (P.VU.vsew) {
case 8: {
require_extension(EXT_ZVFQBDOT8F);
if (P.VU.altfmt) {
ZVBDOT_LOOP(uint8_t, uint8_t, float32_t, zvfqbdot8f_dot_acc<ofp8_e5m2 COMMA ofp8_e4m3>);
} else {
ZVBDOT_LOOP(uint8_t, uint8_t, float32_t, zvfqbdot8f_dot_acc<ofp8_e4m3 COMMA ofp8_e4m3>);
}
break;
}
default: require(false);
}

1
riscv/isa_parser.h

@ -70,6 +70,7 @@ typedef enum {
EXT_ZVQDOTQ,
EXT_ZVQBDOT8I,
EXT_ZVQBDOT16I,
EXT_ZVFQBDOT8F,
EXT_ZVFWBDOT16BF,
EXT_ZVFBDOT32F,
EXT_ZVQLDOT8I,

2
riscv/riscv.mk.in

@ -1079,6 +1079,8 @@ riscv_insn_ext_zvbdot = \
vqbdots_vv \
vfwbdot_vv \
vfbdot_vv \
vfqbdot_vv \
vfqbdot_alt_vv \
riscv_insn_ext_zvldot = \
vqldotu_vv \

2
riscv/vector_unit.cc

@ -46,6 +46,8 @@ reg_t vectorUnit_t::vectorUnit_t::set_vl(int rd, int rs1, reg_t reqVL, reg_t new
ill_altfmt = false;
else if (p->extension_enabled(EXT_ZVQBDOT16I) && vsew == 16)
ill_altfmt = false;
else if (p->extension_enabled(EXT_ZVFQBDOT8F) && vsew == 8)
ill_altfmt = false;
else if (p->extension_enabled(EXT_ZVFWBDOT16BF) && vsew == 16)
ill_altfmt = false;
else if (p->extension_enabled(EXT_ZVQLDOT8I) && vsew == 8)

15
riscv/zvbdot.h

@ -41,4 +41,19 @@ static inline float32_t zvfwbdot16bf_dot_acc(const std::vector<uint16_t>& a, con
return f32_add_odd(f32(res.out), c);
}
template<typename A, typename B>
float32_t zvfqbdot8f_dot_acc(const std::vector<uint8_t>& a, const std::vector<uint8_t>& b, float32_t c)
{
std::vector<A> fa(a.size());
std::transform(a.begin(), a.end(), fa.begin(), [](auto f) { return f; });
std::vector<B> fb(b.size());
std::transform(b.begin(), b.end(), fb.begin(), [](auto f) { return f; });
DotConfig cfg(a.size(), int_log2(a.size()) + ((a.size() & (a.size() - 1)) != 0));
auto res = bulk_norm_dot_ofp8(cfg, &fa[0], &fb[0]);
softfloat_exceptionFlags |= res.flags;
return f32_add_odd(f32(res.out), c);
}
#endif

Loading…
Cancel
Save