#![no_std]
#![no_main]
use et_abi::{DeviceArgs, MINIONS_PER_SHIRE, TensorExtTestArgs};
use et_kernel::{
fence, hart_id, kernel_entry, shire_id,
tensor::{
TensorEvent, fma16a32_xs, ima8a32_xs, tensor_fma16a32, tensor_ima8a32, tensor_load,
tensor_load_b, tensor_store, tensor_store_from_scp, tensor_wait,
},
};
kernel_entry!();
const OUT_STRIDE: usize = 192;
#[unsafe(no_mangle)]
pub extern "C" fn entry_point(args_ptr: usize) -> i64 {
let args: &TensorExtTestArgs = unsafe { TensorExtTestArgs::from_ptr(args_ptr as *const u8) };
let h = hart_id();
if h & 1 != 0 {
return 0;
}
let shire = shire_id();
let minion_in_shire = (h & 63) >> 1;
let my_minion = shire * MINIONS_PER_SHIRE + minion_in_shire;
let total_minions = args.n_shires as u32 * MINIONS_PER_SHIRE;
if my_minion >= total_minions {
return 0;
}
let out_base = args.output as usize + my_minion as usize * OUT_STRIDE;
unsafe { run_tests(args, out_base) };
0
}
#[inline(always)]
unsafe fn run_tests(args: &TensorExtTestArgs, out_base: usize) {
unsafe {
tensor_load(args.a_fp16 as usize, 0, 0, false, 64);
tensor_wait(TensorEvent::Load0);
tensor_load_b(args.b_fp16 as usize, 0, false, 64, true);
tensor_fma16a32(fma16a32_xs(0, 0, 0, 0, true, 0, 0, true, false));
tensor_wait(TensorEvent::Fma);
tensor_store(out_base, 0, 64);
tensor_wait(TensorEvent::Store);
}
unsafe {
tensor_load(args.a_int8 as usize, 1, 0, false, 64);
tensor_load(args.b_int8 as usize, 16, 0, false, 64);
tensor_wait(TensorEvent::Load0);
tensor_ima8a32(ima8a32_xs(
0, 0, 0, 0, false, 16, 1, true, false, false, true, false,
));
tensor_wait(TensorEvent::Fma);
tensor_store(out_base + 64, 0, 64);
tensor_wait(TensorEvent::Store);
}
unsafe {
tensor_store_from_scp(out_base + 128, 0, 0, 1, 64);
tensor_wait(TensorEvent::Store);
}
fence();
}
#[panic_handler]
fn panic(_: &core::panic::PanicInfo) -> ! {
loop {
unsafe { core::arch::asm!("wfi") };
}
}