use crate::exec::compute::vector::context::{VecExecCtx, VecExecResult, mask_active};
use crate::exec::compute::vector::regfile::VectorRegFile;
use crate::isa::fp::FpFlags;
use crate::isa::op::{MaskLogicalOp, MaskOp, MaskSetOp};
use crate::isa::rvv::{ElemIdx, VRegIdx, Vlmax};
pub fn vec_mask_execute(
op: MaskOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
match op {
MaskOp::Logical(op) => exec_mask_logical(op, vpr, vd, vs2, vs1, ctx),
MaskOp::CPop => exec_vcpop(vpr, vs2, ctx),
MaskOp::First => exec_vfirst(vpr, vs2, ctx),
MaskOp::Set(op) => exec_mask_set(op, vpr, vd, vs2, ctx),
MaskOp::Iota => exec_viota(vpr, vd, vs2, ctx),
MaskOp::Id => exec_vid(vpr, vd, ctx),
}
}
#[inline]
const fn no_result() -> VecExecResult {
VecExecResult { vxsat: false, scalar_result: None, fp_flags: FpFlags::NONE }
}
#[inline]
const fn scalar_result(val: u64) -> VecExecResult {
VecExecResult { vxsat: false, scalar_result: Some(val), fp_flags: FpFlags::NONE }
}
#[inline]
const fn compute_mask_logical(op: MaskLogicalOp, s2: bool, s1: bool) -> bool {
match op {
MaskLogicalOp::And => s2 & s1,
MaskLogicalOp::Nand => !(s2 & s1),
MaskLogicalOp::AndNot => s2 && !s1,
MaskLogicalOp::Or => s2 | s1,
MaskLogicalOp::Nor => !(s2 | s1),
MaskLogicalOp::OrNot => s2 || !s1,
MaskLogicalOp::Xor => s2 ^ s1,
MaskLogicalOp::Xnor => !(s2 ^ s1),
}
}
fn exec_mask_logical(
op: MaskLogicalOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
vs1: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let vlen_bits = vpr.vlen().bits();
for i in 0..vlen_bits {
if i < ctx.vstart {
continue;
}
if i >= ctx.vl {
if ctx.vta.is_agnostic() {
vpr.write_mask_bit(vd, ElemIdx::new(i), true);
}
continue;
}
let s2 = vpr.read_mask_bit(vs2, ElemIdx::new(i));
let s1 = vpr.read_mask_bit(vs1, ElemIdx::new(i));
let result = compute_mask_logical(op, s2, s1);
vpr.write_mask_bit(vd, ElemIdx::new(i), result);
}
no_result()
}
fn exec_vcpop(vpr: &impl VectorRegFile, vs2: VRegIdx, ctx: &VecExecCtx) -> VecExecResult {
let mut count: u64 = 0;
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
if vpr.read_mask_bit(vs2, ElemIdx::new(i)) {
count += 1;
}
}
scalar_result(count)
}
fn exec_vfirst(vpr: &impl VectorRegFile, vs2: VRegIdx, ctx: &VecExecCtx) -> VecExecResult {
for i in ctx.vstart..ctx.vl {
if !ctx.vm && !mask_active(vpr, i) {
continue;
}
if vpr.read_mask_bit(vs2, ElemIdx::new(i)) {
return scalar_result(i as u64);
}
}
scalar_result(u64::MAX)
}
fn exec_mask_set(
op: MaskSetOp,
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let vlen_bits = vpr.vlen().bits();
let mut found_first = false;
for i in 0..vlen_bits {
if i < ctx.vstart {
continue;
}
if i >= ctx.vl {
if ctx.vta.is_agnostic() {
vpr.write_mask_bit(vd, ElemIdx::new(i), true);
}
continue;
}
if !ctx.vm && !mask_active(vpr, i) {
if ctx.vma.is_agnostic() {
vpr.write_mask_bit(vd, ElemIdx::new(i), true);
}
continue;
}
let src_bit = vpr.read_mask_bit(vs2, ElemIdx::new(i));
let result = if found_first {
false
} else if src_bit {
found_first = true;
match op {
MaskSetOp::BeforeFirst => false,
MaskSetOp::IncludingFirst | MaskSetOp::OnlyFirst => true,
}
} else {
match op {
MaskSetOp::BeforeFirst | MaskSetOp::IncludingFirst => true,
MaskSetOp::OnlyFirst => false,
}
};
vpr.write_mask_bit(vd, ElemIdx::new(i), result);
}
no_result()
}
fn exec_viota(
vpr: &mut impl VectorRegFile,
vd: VRegIdx,
vs2: VRegIdx,
ctx: &VecExecCtx,
) -> VecExecResult {
let vlmax = Vlmax::compute(vpr.vlen(), ctx.sew, ctx.vlmul).as_usize();
let mut running_sum: u64 = 0;
for i in 0..vlmax {
if i < ctx.vstart {
continue;
}
if i >= ctx.vl {
if ctx.vta.is_agnostic() {
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, ctx.sew.ones());
}
continue;
}
let active = ctx.vm || mask_active(vpr, i);
let prefix = running_sum;
if active && vpr.read_mask_bit(vs2, ElemIdx::new(i)) {
running_sum += 1;
}
if !active {
if ctx.vma.is_agnostic() {
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, ctx.sew.ones());
}
continue;
}
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, prefix);
}
no_result()
}
fn exec_vid(vpr: &mut impl VectorRegFile, vd: VRegIdx, ctx: &VecExecCtx) -> VecExecResult {
let vlmax = Vlmax::compute(vpr.vlen(), ctx.sew, ctx.vlmul).as_usize();
for i in 0..vlmax {
if i < ctx.vstart {
continue;
}
if i >= ctx.vl {
if ctx.vta.is_agnostic() {
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, ctx.sew.ones());
}
continue;
}
if !ctx.vm && !mask_active(vpr, i) {
if ctx.vma.is_agnostic() {
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, ctx.sew.ones());
}
continue;
}
vpr.write_element(vd, ElemIdx::new(i), ctx.sew, i as u64);
}
no_result()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::isa::op::{VecClass, VectorOp};
fn mask_op(op: VectorOp) -> MaskOp {
match op.class() {
VecClass::Mask(mask) => mask,
other => panic!("{op:?} is not a mask op: {other:?}"),
}
}
use crate::arch::regs::vpr::Vpr;
use crate::isa::fp::RoundingMode;
use crate::isa::rvv::{MaskPolicy, Sew, TailPolicy, Vlen, Vlmul, Vxrm};
fn test_vpr() -> Vpr {
Vpr::new(Vlen::new_unchecked(128))
}
fn default_ctx(vl: usize) -> VecExecCtx {
VecExecCtx {
sew: Sew::E32,
vl,
vstart: 0,
vma: MaskPolicy::Undisturbed,
vta: TailPolicy::Undisturbed,
vlmul: Vlmul::M1,
vm: true,
vxrm: Vxrm::RoundToNearestUp,
frm: RoundingMode::Rne,
zvfh: false,
}
}
#[test]
fn test_vid_mf8_e8() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
for i in 0..16usize {
vpr.write_element(vd, ElemIdx::new(i), Sew::E8, 0xAA);
}
let ctx = VecExecCtx {
sew: Sew::E8,
vl: 2,
vstart: 0,
vma: MaskPolicy::Undisturbed,
vta: TailPolicy::Undisturbed,
vlmul: Vlmul::Mf8,
vm: true,
vxrm: Vxrm::RoundToNearestUp,
frm: RoundingMode::Rne,
zvfh: false,
};
let _ = exec_vid(&mut vpr, vd, &ctx);
assert_eq!(vpr.read_element(vd, ElemIdx::new(0), Sew::E8), 0);
assert_eq!(vpr.read_element(vd, ElemIdx::new(1), Sew::E8), 1);
for i in 2..16 {
assert_eq!(
vpr.read_element(vd, ElemIdx::new(i), Sew::E8),
0xAA,
"tail byte {i} should be preserved (tu)"
);
}
}
#[test]
fn test_vmand() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
let vs1 = VRegIdx::new(4);
for i in 0..4 {
vpr.write_mask_bit(vs2, ElemIdx::new(i), true);
}
vpr.write_mask_bit(vs1, ElemIdx::new(0), true);
vpr.write_mask_bit(vs1, ElemIdx::new(2), true);
let ctx = default_ctx(4);
let _ = vec_mask_execute(mask_op(VectorOp::VMAndMM), &mut vpr, vd, vs2, vs1, &ctx);
assert!(vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(1)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(2)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(3)));
}
#[test]
fn test_vmnand() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
let vs1 = VRegIdx::new(4);
vpr.write_mask_bit(vs2, ElemIdx::new(0), true);
vpr.write_mask_bit(vs1, ElemIdx::new(0), true);
let ctx = default_ctx(2);
let _ = vec_mask_execute(mask_op(VectorOp::VMNandMM), &mut vpr, vd, vs2, vs1, &ctx);
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(1)));
}
#[test]
fn test_vmxnor() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
let vs1 = VRegIdx::new(4);
vpr.write_mask_bit(vs2, ElemIdx::new(0), true);
vpr.write_mask_bit(vs1, ElemIdx::new(0), true);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
let ctx = default_ctx(3);
let _ = vec_mask_execute(mask_op(VectorOp::VMXnorMM), &mut vpr, vd, vs2, vs1, &ctx);
assert!(vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(1)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(2)));
}
#[test]
fn test_mask_logical_tail_agnostic() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
let vs1 = VRegIdx::new(4);
vpr.write_mask_bit(vd, ElemIdx::new(2), false);
vpr.write_mask_bit(vd, ElemIdx::new(3), true);
let mut ctx = default_ctx(2);
ctx.vta = TailPolicy::Agnostic;
let _ = vec_mask_execute(mask_op(VectorOp::VMAndMM), &mut vpr, vd, vs2, vs1, &ctx);
assert!(vpr.read_mask_bit(vd, ElemIdx::new(2)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(3)));
}
#[test]
fn test_vcpop_unmasked() {
let mut vpr = test_vpr();
let vs2 = VRegIdx::new(1);
vpr.write_mask_bit(vs2, ElemIdx::new(0), true);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
vpr.write_mask_bit(vs2, ElemIdx::new(3), true);
let ctx = default_ctx(4);
let result = exec_vcpop(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(3));
}
#[test]
fn test_vcpop_masked() {
let mut vpr = test_vpr();
let vs2 = VRegIdx::new(1);
let v0 = VRegIdx::new(0);
for i in 0..3 {
vpr.write_mask_bit(vs2, ElemIdx::new(i), true);
}
vpr.write_mask_bit(v0, ElemIdx::new(0), true);
vpr.write_mask_bit(v0, ElemIdx::new(2), true);
let mut ctx = default_ctx(3);
ctx.vm = false;
let result = exec_vcpop(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(2));
}
#[test]
fn test_vcpop_with_vstart() {
let mut vpr = test_vpr();
let vs2 = VRegIdx::new(1);
for i in 0..4 {
vpr.write_mask_bit(vs2, ElemIdx::new(i), true);
}
let mut ctx = default_ctx(4);
ctx.vstart = 2;
let result = exec_vcpop(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(2));
}
#[test]
fn test_vfirst_found() {
let mut vpr = test_vpr();
let vs2 = VRegIdx::new(1);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
let ctx = default_ctx(4);
let result = exec_vfirst(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(2));
}
#[test]
fn test_vfirst_not_found() {
let vpr = test_vpr();
let vs2 = VRegIdx::new(1);
let ctx = default_ctx(4);
let result = exec_vfirst(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(u64::MAX));
}
#[test]
fn test_vfirst_masked() {
let mut vpr = test_vpr();
let vs2 = VRegIdx::new(1);
let v0 = VRegIdx::new(0);
vpr.write_mask_bit(vs2, ElemIdx::new(0), true);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
vpr.write_mask_bit(v0, ElemIdx::new(2), true);
let mut ctx = default_ctx(4);
ctx.vm = false;
let result = exec_vfirst(&vpr, vs2, &ctx);
assert_eq!(result.scalar_result, Some(2));
}
#[test]
fn test_vmsbf() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
vpr.write_mask_bit(vs2, ElemIdx::new(3), true);
let ctx = default_ctx(4);
let _ = exec_mask_set(MaskSetOp::BeforeFirst, &mut vpr, vd, vs2, &ctx);
assert!(vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(1)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(2)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(3)));
}
#[test]
fn test_vmsif() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
let ctx = default_ctx(4);
let _ = exec_mask_set(MaskSetOp::IncludingFirst, &mut vpr, vd, vs2, &ctx);
assert!(vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(1)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(2)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(3)));
}
#[test]
fn test_vmsof() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
vpr.write_mask_bit(vs2, ElemIdx::new(3), true);
let ctx = default_ctx(4);
let _ = exec_mask_set(MaskSetOp::OnlyFirst, &mut vpr, vd, vs2, &ctx);
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(0)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(1)));
assert!(vpr.read_mask_bit(vd, ElemIdx::new(2)));
assert!(!vpr.read_mask_bit(vd, ElemIdx::new(3)));
}
#[test]
fn test_vmsbf_no_set_bit() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
let ctx = default_ctx(4);
let _ = exec_mask_set(MaskSetOp::BeforeFirst, &mut vpr, vd, vs2, &ctx);
for i in 0..4 {
assert!(vpr.read_mask_bit(vd, ElemIdx::new(i)));
}
}
#[test]
fn test_viota_basic() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let vs2 = VRegIdx::new(3);
vpr.write_mask_bit(vs2, ElemIdx::new(0), true);
vpr.write_mask_bit(vs2, ElemIdx::new(2), true);
let ctx = default_ctx(4);
let _ = exec_viota(&mut vpr, vd, vs2, &ctx);
assert_eq!(vpr.read_element(vd, ElemIdx::new(0), Sew::E32), 0);
assert_eq!(vpr.read_element(vd, ElemIdx::new(1), Sew::E32), 1);
assert_eq!(vpr.read_element(vd, ElemIdx::new(2), Sew::E32), 1);
assert_eq!(vpr.read_element(vd, ElemIdx::new(3), Sew::E32), 2);
}
#[test]
fn test_vid_basic() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let ctx = default_ctx(4);
let _ = exec_vid(&mut vpr, vd, &ctx);
for i in 0..4u64 {
assert_eq!(vpr.read_element(vd, ElemIdx::new(i as usize), Sew::E32), i);
}
}
#[test]
fn test_vid_masked() {
let mut vpr = test_vpr();
let vd = VRegIdx::new(2);
let v0 = VRegIdx::new(0);
for i in 0..4 {
vpr.write_element(vd, ElemIdx::new(i), Sew::E32, 0xFF);
}
vpr.write_mask_bit(v0, ElemIdx::new(0), true);
vpr.write_mask_bit(v0, ElemIdx::new(2), true);
let mut ctx = default_ctx(4);
ctx.vm = false;
let _ = exec_vid(&mut vpr, vd, &ctx);
assert_eq!(vpr.read_element(vd, ElemIdx::new(0), Sew::E32), 0);
assert_eq!(vpr.read_element(vd, ElemIdx::new(1), Sew::E32), 0xFF); assert_eq!(vpr.read_element(vd, ElemIdx::new(2), Sew::E32), 2);
assert_eq!(vpr.read_element(vd, ElemIdx::new(3), Sew::E32), 0xFF); }
}