#![no_std]
#![no_main]
use et_abi::{
DeviceArgs, GemmArgs, GEMM_TILE_K, GEMM_TILE_M, GEMM_TILE_N, MINIONS_PER_SHIRE,
};
use et_kernel::{
fence, hart_id, kernel_entry, shire_id,
tensor::{
TensorEvent, fma32_xs, tensor_fma32, tensor_load, tensor_load_b,
tensor_store, tensor_wait,
},
};
kernel_entry!();
#[unsafe(no_mangle)]
pub extern "C" fn entry_point(args_ptr: usize) -> i64 {
let args: &GemmArgs = unsafe { GemmArgs::from_ptr(args_ptr as *const u8) };
let h = hart_id();
if h & 1 != 0 {
return 0;
}
let shire = shire_id();
let hart_in_shire = h & 63; let minion_in_shire = hart_in_shire >> 1;
let total_minions = args.n_shires as u32 * MINIONS_PER_SHIRE;
let my_minion = shire * MINIONS_PER_SHIRE + minion_in_shire;
if my_minion >= total_minions {
return 0;
}
let n_tile_m = (args.m as usize).div_ceil(GEMM_TILE_M);
let n_tile_n = (args.n as usize).div_ceil(GEMM_TILE_N);
let n_tiles = n_tile_m * n_tile_n;
let shire_size = n_tiles.div_ceil(args.n_shires as usize);
let shire_base = (shire as usize) * shire_size;
let shire_end = n_tiles.saturating_sub(shire_base).min(shire_size);
let mut local_idx = minion_in_shire as usize;
while local_idx < shire_end {
let tile_idx = shire_base + local_idx;
let tile_row = tile_idx / n_tile_n;
let tile_col = tile_idx % n_tile_n;
unsafe { compute_tile(args, tile_row, tile_col) };
local_idx += MINIONS_PER_SHIRE as usize;
}
0
}
#[inline(always)]
unsafe fn compute_tile(args: &GemmArgs, tile_row: usize, tile_col: usize) {
let c_row = tile_row * GEMM_TILE_M;
let c_col = tile_col * GEMM_TILE_N;
let actual_m = GEMM_TILE_M.min(args.m as usize - c_row);
let arows = (actual_m - 1) as u8;
let bcols: u8 = 3;
let a_base = args.a as usize;
let b_base = args.b as usize;
let c_base = args.c as usize;
let lda = args.lda as usize; let ldb = args.ldb as usize;
let ldc = args.ldc as usize;
let n_k_tiles = (args.k as usize).div_ceil(GEMM_TILE_K);
let mut k_tile = 0_usize;
while k_tile < n_k_tiles {
let k_start = k_tile * GEMM_TILE_K;
let actual_k = GEMM_TILE_K.min(args.k as usize - k_start);
let acols = (actual_k - 1) as u8;
let a_addr = a_base + c_row * lda + k_start * 4;
let b_addr = b_base + k_start * ldb + c_col * 4;
unsafe {
tensor_load(a_addr, 0, arows, false, lda as u64);
tensor_wait(TensorEvent::Load0);
}
unsafe {
tensor_load_b(b_addr, acols, false, ldb as u64, true);
}
let xs = fma32_xs(
bcols,
arows,
acols,
0, true, 0, 0, k_tile == 0, false, );
unsafe {
tensor_fma32(xs);
tensor_wait(TensorEvent::Fma);
}
k_tile += 1;
}
let c_addr = c_base + c_row * ldc + c_col * 4;
unsafe {
tensor_store(c_addr, arows, ldc as u64);
}
fence();
}
#[panic_handler]
fn panic(_info: &core::panic::PanicInfo) -> ! {
loop {
unsafe { core::arch::asm!("wfi") };
}
}