#![allow(clippy::float_cmp)]
use crate::exec::compute::fpu::half::{CANONICAL_NAN_F16, f16_to_f32, f64_to_f16, is_snan_f16};
use crate::exec::compute::fpu::host::{
clear_host_fp_flags, read_host_fp_flags, restore_host_round_mode, set_host_round_mode,
};
use crate::exec::compute::fpu::nan_handling::{
box_f32_canon, canonicalize_f64_bits, fmax_f32, fmax_f64, fmin_f32, fmin_f64,
};
use crate::exec::compute::fpu::nan_handling::{is_snan_f32, is_snan_f64};
use crate::exec::compute::vector::context::{
FpSew, FpWiden, VecExecCtx, VecExecResult, mask_active, sign_extend, widen_sew,
};
use crate::exec::compute::vector::regfile::VectorRegFile;
use crate::isa::fp::{FpFlags, RoundingMode};
use crate::isa::op::{
FpReduceOp, FpWidenReduceOp, IntReduceOp, ReduceOp, VectorOp, WidenIntReduceOp,
};
use crate::isa::rvv::{ElemIdx, Sew, VRegIdx, Vlmax, Vlmul};
pub const fn is_reduction(op: VectorOp) -> bool {
matches!(
op,
VectorOp::VRedSum
| VectorOp::VRedAnd
| VectorOp::VRedOr
| VectorOp::VRedXor
| VectorOp::VRedMinU
| VectorOp::VRedMin
| VectorOp::VRedMaxU
| VectorOp::VRedMax
| VectorOp::VWRedSumU
| VectorOp::VWRedSum
| VectorOp::VFRedOSum
| VectorOp::VFRedUSum
| VectorOp::VFRedMax
| VectorOp::VFRedMin
| VectorOp::VFWRedOSum
| VectorOp::VFWRedUSum
)
}
pub fn vec_reduce(
op: ReduceOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
match op {
ReduceOp::Int(op) => exec_int_reduction(op, vpr, vd, vs2, vs1, ctx),
ReduceOp::WidenInt(op) => exec_widen_int_reduction(op, vpr, vd, vs2, vs1, ctx),
ReduceOp::Fp(op) => exec_fp_reduction(op, vpr, vd, vs2, vs1, ctx),
ReduceOp::FpWiden(op) => exec_fp_widen_reduction(op, vpr, vd, vs2, vs1, ctx),
}
}
#[inline]
fn read_initial_accum(vpr: &impl VectorRegFile, vs1: VRegIdx, sew: Sew) -> u64 {
vpr.read_element(vs1, ElemIdx::new(0), sew)
}
fn exec_int_reduction(
op: IntReduceOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let sew = ctx.sew;
let mask = sew.mask();
let mut acc = read_initial_accum(vpr, vs1, sew);
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem = vpr.read_element(vs2, ElemIdx::new(i), sew);
acc = int_reduce_step(op, acc, elem, sew, mask);
}
vpr.write_element(vd, ElemIdx::new(0), sew, acc & mask);
if ctx.vta.is_agnostic() {
let vlmax = Vlmax::compute(vpr.vlen(), sew, Vlmul::M1).as_usize();
for i in 1..vlmax {
vpr.write_element(vd, ElemIdx::new(i), sew, sew.ones());
}
}
VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE }
}
#[inline]
const fn int_reduce_step(op: IntReduceOp, acc: u64, elem: u64, sew: Sew, mask: u64) -> u64 {
match op {
IntReduceOp::Sum => acc.wrapping_add(elem) & mask,
IntReduceOp::And => acc & elem,
IntReduceOp::Or => acc | elem,
IntReduceOp::Xor => acc ^ elem,
IntReduceOp::MinU => {
if elem < acc {
elem
} else {
acc
}
}
IntReduceOp::Min => {
let sa = sign_extend(acc, sew);
let se = sign_extend(elem, sew);
if se < sa { elem } else { acc }
}
IntReduceOp::MaxU => {
if elem > acc {
elem
} else {
acc
}
}
IntReduceOp::Max => {
let sa = sign_extend(acc, sew);
let se = sign_extend(elem, sew);
if se > sa { elem } else { acc }
}
}
}
fn exec_widen_int_reduction(
op: WidenIntReduceOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let src_sew = ctx.sew;
let Some(dst_sew) = widen_sew(src_sew) else {
return VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE };
};
let dst_mask = dst_sew.mask();
let mut acc = read_initial_accum(vpr, vs1, dst_sew);
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem = vpr.read_element(vs2, ElemIdx::new(i), src_sew);
let wide = match op {
WidenIntReduceOp::SumU => elem,
WidenIntReduceOp::Sum => (sign_extend(elem, src_sew) as u64) & dst_mask,
};
acc = acc.wrapping_add(wide) & dst_mask;
}
vpr.write_element(vd, ElemIdx::new(0), dst_sew, acc);
if ctx.vta.is_agnostic() {
let vlmax = Vlmax::compute(vpr.vlen(), dst_sew, Vlmul::M1).as_usize();
for i in 1..vlmax {
vpr.write_element(vd, ElemIdx::new(i), dst_sew, dst_sew.ones());
}
}
VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE }
}
fn exec_fp_reduction(
op: FpReduceOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let sew = ctx.sew;
let Some(fp_sew) = FpSew::of(sew, ctx.zvfh) else {
return VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE };
};
let saved_rm = set_host_round_mode(ctx.frm);
let (result_bits, flags) = match fp_sew {
FpSew::F16 => fp_reduce_f16(op, vpr, vs2, vs1, ctx),
FpSew::F32 => fp_reduce_f32(op, vpr, vs2, vs1, ctx),
FpSew::F64 => fp_reduce_f64(op, vpr, vs2, vs1, ctx),
};
vpr.write_element(vd, ElemIdx::new(0), sew, result_bits);
restore_host_round_mode(saved_rm);
if ctx.vta.is_agnostic() {
let vlmax = Vlmax::compute(vpr.vlen(), sew, Vlmul::M1).as_usize();
for i in 1..vlmax {
vpr.write_element(vd, ElemIdx::new(i), sew, sew.ones());
}
}
VecExecResult { vxsat: false, scalar_result: None, fp_flags: flags }
}
fn fp_reduce_f32(
op: FpReduceOp,
vpr: &impl VectorRegFile,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> (u64, FpFlags) {
let sew = ctx.sew;
let init_bits = read_initial_accum(vpr, vs1, sew);
let mut acc = f32::from_bits(init_bits as u32);
let mut flags = FpFlags::NONE;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem_bits = vpr.read_element(vs2, ElemIdx::new(i), sew);
let elem = f32::from_bits(elem_bits as u32);
match op {
FpReduceOp::OSum | FpReduceOp::USum => {
clear_host_fp_flags();
acc = std::hint::black_box(std::hint::black_box(acc) + std::hint::black_box(elem));
flags = flags | read_host_fp_flags();
}
FpReduceOp::Min => {
if is_snan_f32(acc) || is_snan_f32(elem) {
flags = flags | FpFlags::NV;
}
acc = fmin_f32(acc, elem);
}
FpReduceOp::Max => {
if is_snan_f32(acc) || is_snan_f32(elem) {
flags = flags | FpFlags::NV;
}
acc = fmax_f32(acc, elem);
}
}
}
(box_f32_canon(acc), flags)
}
fn fp_reduce_f64(
op: FpReduceOp,
vpr: &impl VectorRegFile,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> (u64, FpFlags) {
let sew = ctx.sew;
let init_bits = read_initial_accum(vpr, vs1, sew);
let mut acc = f64::from_bits(init_bits);
let mut flags = FpFlags::NONE;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem_bits = vpr.read_element(vs2, ElemIdx::new(i), sew);
let elem = f64::from_bits(elem_bits);
match op {
FpReduceOp::OSum | FpReduceOp::USum => {
clear_host_fp_flags();
acc = std::hint::black_box(std::hint::black_box(acc) + std::hint::black_box(elem));
flags = flags | read_host_fp_flags();
}
FpReduceOp::Min => {
if is_snan_f64(acc) || is_snan_f64(elem) {
flags = flags | FpFlags::NV;
}
acc = fmin_f64(acc, elem);
}
FpReduceOp::Max => {
if is_snan_f64(acc) || is_snan_f64(elem) {
flags = flags | FpFlags::NV;
}
acc = fmax_f64(acc, elem);
}
}
}
(canonicalize_f64_bits(acc), flags)
}
fn fp_reduce_f16(
op: FpReduceOp,
vpr: &impl VectorRegFile,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> (u64, FpFlags) {
let sew = ctx.sew;
let init_bits = read_initial_accum(vpr, vs1, sew) as u16;
let mut acc_bits: u16 = init_bits;
let mut flags = FpFlags::NONE;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem_bits = vpr.read_element(vs2, ElemIdx::new(i), sew) as u16;
match op {
FpReduceOp::OSum | FpReduceOp::USum => {
if is_snan_f16(acc_bits) || is_snan_f16(elem_bits) {
flags = flags | FpFlags::NV;
}
let acc_f64 = f16_to_f32(acc_bits) as f64;
let elem_f64 = f16_to_f32(elem_bits) as f64;
clear_host_fp_flags();
let sum = std::hint::black_box(
std::hint::black_box(acc_f64) + std::hint::black_box(elem_f64),
);
let host_flags = read_host_fp_flags();
let (rounded, round_flags) = f64_to_f16(sum, ctx.frm);
acc_bits = rounded;
flags = flags | host_flags | round_flags;
}
FpReduceOp::Min => {
if is_snan_f16(acc_bits) || is_snan_f16(elem_bits) {
flags = flags | FpFlags::NV;
}
let a = f16_to_f32(acc_bits);
let b = f16_to_f32(elem_bits);
let r = fmin_f32(a, b);
acc_bits = if r.is_nan() {
CANONICAL_NAN_F16
} else {
let (bits, _) = f64_to_f16(r as f64, RoundingMode::Rne);
bits
};
}
FpReduceOp::Max => {
if is_snan_f16(acc_bits) || is_snan_f16(elem_bits) {
flags = flags | FpFlags::NV;
}
let a = f16_to_f32(acc_bits);
let b = f16_to_f32(elem_bits);
let r = fmax_f32(a, b);
acc_bits = if r.is_nan() {
CANONICAL_NAN_F16
} else {
let (bits, _) = f64_to_f16(r as f64, RoundingMode::Rne);
bits
};
}
}
}
(acc_bits as u64, flags)
}
fn exec_fp_widen_reduction(
op: FpWidenReduceOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let Some(widen) = FpSew::of(ctx.sew, ctx.zvfh).and_then(FpSew::widening) else {
return VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE };
};
let dst_sew = widen.dst();
let saved_rm = set_host_round_mode(ctx.frm);
let (result_bits, flags) = match widen {
FpWiden::F32ToF64 => fp_widen_reduce_f32_to_f64(op, vpr, vs2, vs1, ctx, dst_sew),
FpWiden::F16ToF32 => fp_widen_reduce_f16_to_f32(op, vpr, vs2, vs1, ctx, dst_sew),
};
vpr.write_element(vd, ElemIdx::new(0), dst_sew, result_bits);
restore_host_round_mode(saved_rm);
if ctx.vta.is_agnostic() {
let vlmax = Vlmax::compute(vpr.vlen(), dst_sew, Vlmul::M1).as_usize();
for i in 1..vlmax {
vpr.write_element(vd, ElemIdx::new(i), dst_sew, dst_sew.ones());
}
}
VecExecResult { vxsat: false, scalar_result: None, fp_flags: flags }
}
fn fp_widen_reduce_f32_to_f64(
op: FpWidenReduceOp,
vpr: &impl VectorRegFile,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
dst_sew: Sew,
) -> (u64, FpFlags) {
let init_bits = read_initial_accum(vpr, vs1, dst_sew);
let mut acc = f64::from_bits(init_bits);
let mut flags = FpFlags::NONE;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem_bits = vpr.read_element(vs2, ElemIdx::new(i), Sew::E32) as u32;
let wide = f32::from_bits(elem_bits) as f64;
match op {
FpWidenReduceOp::OSum | FpWidenReduceOp::USum => {
clear_host_fp_flags();
acc = std::hint::black_box(std::hint::black_box(acc) + std::hint::black_box(wide));
flags = flags | read_host_fp_flags();
}
}
}
(canonicalize_f64_bits(acc), flags)
}
fn fp_widen_reduce_f16_to_f32(
op: FpWidenReduceOp,
vpr: &impl VectorRegFile,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
dst_sew: Sew,
) -> (u64, FpFlags) {
let init_bits = read_initial_accum(vpr, vs1, dst_sew) as u32;
let mut acc = f32::from_bits(init_bits);
let mut flags = FpFlags::NONE;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
let elem_bits = vpr.read_element(vs2, ElemIdx::new(i), Sew::E16) as u16;
let wide = f16_to_f32(elem_bits);
match op {
FpWidenReduceOp::OSum | FpWidenReduceOp::USum => {
if is_snan_f16(elem_bits) {
flags = flags | FpFlags::NV;
}
clear_host_fp_flags();
acc = std::hint::black_box(std::hint::black_box(acc) + std::hint::black_box(wide));
flags = flags | read_host_fp_flags();
}
}
}
(box_f32_canon(acc), flags)
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
use crate::isa::op::VecClass;
fn reduce_op(op: VectorOp) -> ReduceOp {
match op.class() {
VecClass::Reduce(reduce) => reduce,
other => panic!("{op:?} is not a reduction: {other:?}"),
}
}
use crate::arch::regs::vpr::Vpr;
use crate::isa::fp::RoundingMode;
use crate::isa::rvv::{MaskPolicy, TailPolicy, Vlen, Vlmul, Vxrm};
fn make_ctx(sew: Sew, vl: usize) -> VecExecCtx {
VecExecCtx {
sew,
vl,
vstart: 0,
vma: MaskPolicy::Undisturbed,
vta: TailPolicy::Undisturbed,
vlmul: Vlmul::M1,
vm: true,
vxrm: Vxrm::RoundToNearestUp,
frm: RoundingMode::Rne,
zvfh: false,
}
}
fn vpr128() -> Vpr {
Vpr::new(Vlen::new_unchecked(128))
}
#[test]
fn test_vredsum() {
let mut vpr = vpr128();
let ctx = make_ctx(Sew::E32, 4);
let vd = VRegIdx::new(1);
let vs2 = VRegIdx::new(2);
let vs1 = VRegIdx::new(3);
for i in 0..4 {
vpr.write_element(vs2, ElemIdx::new(i), Sew::E32, (i as u64 + 1) * 10);
}
vpr.write_element(vs1, ElemIdx::new(0), Sew::E32, 100);
let _result = vec_reduce(reduce_op(VectorOp::VRedSum), &mut vpr, vd, vs2, vs1, &ctx);
assert_eq!(vpr.read_element(vd, ElemIdx::new(0), Sew::E32), 200);
}
#[test]
fn test_vredand() {
let mut vpr = vpr128();
let ctx = make_ctx(Sew::E32, 4);
let vd = VRegIdx::new(1);
let vs2 = VRegIdx::new(2);
let vs1 = VRegIdx::new(3);
for i in 0..4 {
vpr.write_element(vs2, ElemIdx::new(i), Sew::E32, 0xFF);
}
vpr.write_element(vs1, ElemIdx::new(0), Sew::E32, 0xFFFF_FFFF);
let _result = vec_reduce(reduce_op(VectorOp::VRedAnd), &mut vpr, vd, vs2, vs1, &ctx);
assert_eq!(vpr.read_element(vd, ElemIdx::new(0), Sew::E32), 0xFF);
}
#[test]
fn test_vredmin_signed() {
let mut vpr = vpr128();
let ctx = make_ctx(Sew::E32, 4);
let vd = VRegIdx::new(1);
let vs2 = VRegIdx::new(2);
let vs1 = VRegIdx::new(3);
vpr.write_element(vs2, ElemIdx::new(0), Sew::E32, 5);
vpr.write_element(vs2, ElemIdx::new(1), Sew::E32, (-3i32 as u32) as u64);
vpr.write_element(vs2, ElemIdx::new(2), Sew::E32, 10);
vpr.write_element(vs2, ElemIdx::new(3), Sew::E32, 1);
vpr.write_element(vs1, ElemIdx::new(0), Sew::E32, 100);
let _result = vec_reduce(reduce_op(VectorOp::VRedMin), &mut vpr, vd, vs2, vs1, &ctx);
let val = vpr.read_element(vd, ElemIdx::new(0), Sew::E32);
assert_eq!(val as u32, (-3i32) as u32);
}
#[test]
fn test_vfredosum_f32() {
let mut vpr = vpr128();
let ctx = make_ctx(Sew::E32, 4);
let vd = VRegIdx::new(1);
let vs2 = VRegIdx::new(2);
let vs1 = VRegIdx::new(3);
for i in 0..4 {
vpr.write_element(vs2, ElemIdx::new(i), Sew::E32, ((i as f32 + 1.0).to_bits()) as u64);
}
vpr.write_element(vs1, ElemIdx::new(0), Sew::E32, 10.0f32.to_bits() as u64);
let _result = vec_reduce(reduce_op(VectorOp::VFRedOSum), &mut vpr, vd, vs2, vs1, &ctx);
let val = f32::from_bits(vpr.read_element(vd, ElemIdx::new(0), Sew::E32) as u32);
assert_eq!(val, 20.0); }
}