diff --git a/.github/workflows/transactional_emulator.yml b/.github/workflows/transactional_emulator.yml index e63788de..141d440b 100644 --- a/.github/workflows/transactional_emulator.yml +++ b/.github/workflows/transactional_emulator.yml @@ -6,6 +6,9 @@ on: branches: [main] paths: - 'transactional_emulator/**' + - 'plena_settings.toml' + - 'justfile' + - '.github/workflows/transactional_emulator.yml' - 'PLENA_Tools/**' # A compiler pin bump is what moves MOE_STAGES, so it has to reach the # stage-attribution guard below. @@ -18,6 +21,9 @@ on: branches: [main] paths: - 'transactional_emulator/**' + - 'plena_settings.toml' + - 'justfile' + - '.github/workflows/transactional_emulator.yml' - 'PLENA_Tools/**' - 'PLENA_Compiler' - 'flake.nix' @@ -34,6 +40,8 @@ jobs: runs-on: ubuntu-latest steps: - uses: actions/checkout@v4 + with: + submodules: recursive # Pinned: with `-D warnings`, a rolling channel reds unrelated PRs. - name: Install Rust toolchain @@ -118,7 +126,7 @@ jobs: - name: Run Rust tests run: | - nix develop --command bash -c "cd transactional_emulator && cargo test --workspace --release" + nix develop --command bash -c "cd transactional_emulator && cargo test --workspace --release -- --test-threads=1" - name: Install uv uses: astral-sh/setup-uv@v4 @@ -147,3 +155,7 @@ jobs: - name: Run emulator timing gates run: | nix develop --command bash -c "just test-timing-gates" + + - name: Run Matrix SRAM Compiler-to-Rust integration + run: | + nix develop --no-write-lock-file --command just test-matrix-lcompute diff --git a/PLENA_Compiler b/PLENA_Compiler index d89ad594..9637858a 160000 --- a/PLENA_Compiler +++ b/PLENA_Compiler @@ -1 +1 @@ -Subproject commit d89ad594c798daa54f63f914aebad5317653489a +Subproject commit 9637858ad9b64ac8c02adcc0ec0b1d749ae6a6fa diff --git a/justfile b/justfile index 8050b109..4aa73d4e 100644 --- a/justfile +++ b/justfile @@ -256,3 +256,6 @@ multilayer-decoder-profile model="smolvlm2": test-sliced-aten-emulator model="AICrossSim/clm-60m" seq_len="64" num_layers="1": cd PLENA_Compiler && PYTHONPATH=".:../PLENA_Tools:../transactional_emulator/testbench:..:" python3 -m compiler.aten.sliced_emulator_runner {{model}} --seq-len {{seq_len}} --num-layers {{num_layers}} +# Matrix SRAM views and prepared BF16 recurrence through Compiler -> Rust. +test-matrix-lcompute *args: + python3 transactional_emulator/testbench/aten/matrix_lcompute_test.py {{args}} diff --git a/plena_settings.toml b/plena_settings.toml index 05fddad9..b6691064 100644 --- a/plena_settings.toml +++ b/plena_settings.toml @@ -325,6 +325,15 @@ sign = false exponent = 8 mantissa = 0 +# Explicit Matrix-view state transfers are independent of activation/KV precision. +[TRANSACTIONAL.PRECISION.HBM_STATE_TYPE] +format = "Plain" +[TRANSACTIONAL.PRECISION.HBM_STATE_TYPE.DATA_TYPE] +type = "Fp" +sign = true +exponent = 8 +mantissa = 7 + [TRANSACTIONAL.PRECISION.HBM_V_INT_TYPE] format = "Plain" [TRANSACTIONAL.PRECISION.HBM_V_INT_TYPE.DATA_TYPE] diff --git a/transactional_emulator/lib/quantize/src/dtype.rs b/transactional_emulator/lib/quantize/src/dtype.rs index a84b299b..862b3d6d 100644 --- a/transactional_emulator/lib/quantize/src/dtype.rs +++ b/transactional_emulator/lib/quantize/src/dtype.rs @@ -428,6 +428,14 @@ impl DataType { /// Convert bytes to vector of f32. pub fn convert_bytes_to_f32_vec(self, mut bytes: &[u8], out: &mut [f32]) { let bits = self.size_in_bits(); + if bits == 32 { + assert!(bytes.len() >= out.len() * 4); + for (encoded, decoded) in bytes.chunks_exact(4).zip(out.iter_mut()) { + let word = u32::from_le_bytes(encoded.try_into().expect("four-byte word")); + *decoded = self.convert_bits_to_f32(word); + } + return; + } let mut data = 0; let mut bits_left = 0; for out in out.iter_mut() { @@ -445,6 +453,13 @@ impl DataType { pub fn bytes_from_f32(self, input: &[f32], mut out: &mut [u8]) { let bits = self.size_in_bits(); + if bits == 32 { + assert!(out.len() >= input.len() * 4); + for (value, encoded) in input.iter().zip(out.chunks_exact_mut(4)) { + encoded.copy_from_slice(&self.bits_from_f32(*value).to_le_bytes()); + } + return; + } let mut data = 0; let mut bits_left = 0u8; @@ -579,6 +594,17 @@ mod tests { assert_eq!(out, vec![1.0, 2.0, 3.0, 4.0]); } + #[test] + fn test_datatype_bytes_roundtrip_32bit_words() { + let ty = DataType::Fp(FpType::F32); + let input = [0.0, 1.0, 2.0, 3.0]; + let mut bytes = vec![0u8; input.len() * 4]; + ty.bytes_from_f32(&input, &mut bytes); + let mut out = vec![0.0; input.len()]; + ty.convert_bytes_to_f32_vec(&bytes, &mut out); + assert_eq!(out, input); + } + #[test] fn test_datatype_size_in_bytes_current_behavior() { // Pinned as-is: this returns `size_in_bits` (it is not divided by 8). diff --git a/transactional_emulator/lib/sram/src/lib.rs b/transactional_emulator/lib/sram/src/lib.rs index 8fb85c08..efeae78e 100644 --- a/transactional_emulator/lib/sram/src/lib.rs +++ b/transactional_emulator/lib/sram/src/lib.rs @@ -48,14 +48,6 @@ impl Cell { } } -impl Cell { - /// Convenience for cells whose payload is already a [`QuantTensor`]; no - /// conversion needed. - pub(crate) async fn resolve(&mut self) -> &QuantTensor { - self.resolve_with(|t| t).await - } -} - /// Convert an element-address into a cell index, asserting both alignment /// (`addr` is a multiple of `units_per_cell`) and bounds (`idx < depth`). pub(crate) fn addr_to_cell(addr: u32, units_per_cell: u32, depth: usize) -> usize { diff --git a/transactional_emulator/lib/sram/src/matrix.rs b/transactional_emulator/lib/sram/src/matrix.rs index 7b64adf3..d388f554 100644 --- a/transactional_emulator/lib/sram/src/matrix.rs +++ b/transactional_emulator/lib/sram/src/matrix.rs @@ -1,33 +1,175 @@ -use quantize::{tensor_to_f32_vec, MxDataType, QuantTensor}; -use tokio::sync::oneshot::Receiver; +use quantize::{tensor_from_f32_slice, tensor_to_f32_vec, DataType, MxDataType, QuantTensor}; +use std::collections::{hash_map::Entry, HashMap}; +use std::sync::atomic::{AtomicU64, Ordering}; +use tokio::sync::oneshot::{Receiver, Sender}; use tokio::sync::Mutex; use crate::{addr_to_cell, Cell}; -/// Behaviour modelling of matrix SRAM. +/// Logical Matrix-SRAM view interpreted by the physical bank mapper. /// -/// The timing aspect is to be considered by the matrix machine itself. +/// `tile_pitch_rows` is measured in physical rows inside each bank. The +/// arithmetic operation is deliberately absent: this structure describes +/// placement only. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct MatrixLayout { + pub rows: u32, + pub cols: u32, + pub tile_count: u32, + pub tile_pitch_rows: u32, + /// Bank skew used by the physical machine. Public Matrix views fix this to 1; + /// tests may vary it only for a non-architectural upper-bound control. + pub alpha: u32, + /// Additional compiler-selected phase between logical tiles. + pub tile_skew: u32, +} + +/// Logical direction serviced from the same physical Matrix-SRAM cells. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MatrixAccessAxis { + Row, + Column, +} + +/// One physical location in a banked Matrix SRAM. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct MatrixPhysicalCoord { + pub bank: u32, + pub bank_row: u32, + pub lane: u32, +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct MatrixPacketService { + pub values: u64, + pub bank_words: u64, + pub ideal_cycles: u64, + pub service_cycles: u64, + pub bank_stall_cycles: u64, + pub worst_bank_words: u64, +} + +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct MatrixPacketCounterSnapshot { + pub packets: u64, + pub values: u64, + pub bank_words: u64, + pub ideal_cycles: u64, + pub service_cycles: u64, + pub bank_stall_cycles: u64, +} + +#[derive(Default)] +struct MatrixPacketCounters { + packets: AtomicU64, + values: AtomicU64, + bank_words: AtomicU64, + ideal_cycles: AtomicU64, + service_cycles: AtomicU64, + bank_stall_cycles: AtomicU64, +} + +struct PendingWord { + sender: Sender, + source_offset: usize, +} + +/// Opaque completion handle returned when a Matrix DMA parks physical words. +pub struct PendingMatrixTiles { + words: Vec, +} + +/// Physically banked Matrix SRAM. +/// +/// A cell is one `bank_width`-element bank word, not a whole matrix tile. A +/// logical tile is reconstructed through the fixed-diagonal placement function on +/// read. Consequently a wrong view changes values, not merely a timing +/// counter, which is the evidence needed for lane-restoration correctness. pub struct MatrixSram { tile_size: u32, - tiles: Vec>>, + /// Number of MLEN-wide logical rows physically implemented. + depth_rows: usize, + /// Complete square tiles addressable by the legacy Matrix APIs. + full_tile_count: usize, ty: MxDataType, + element_type: DataType, + banks: u32, + bank_width: u32, + rows_per_tile: u32, + bank_rows: Vec>>>>, + packet_counters: MatrixPacketCounters, } impl MatrixSram { - /// Create a matrix SRAM with given tile size and depth. + /// Backwards-compatible constructor with at most 64 physical banks. + /// + /// Small test geometries retain one scalar lane per bank. Wider, paper- + /// scale rows widen each bank word instead of exceeding the six-bit bank + /// index used by Matrix views. pub fn new(tile_size: u32, depth: usize, ty: MxDataType) -> Self { - let tiles = (0..(depth / tile_size as usize)) + assert!( + tile_size.is_power_of_two(), + "legacy Matrix tile size must be a power of two" + ); + let bank_width = (tile_size / 64).max(1); + Self::with_banks(tile_size, depth, bank_width, ty) + } + + /// Construct the Matrix SRAM with fixed diagonal placement and + /// `tile_size / bank_width` physical banks. + pub fn with_banks(tile_size: u32, depth: usize, bank_width: u32, ty: MxDataType) -> Self { + assert!(tile_size > 0); + assert!(bank_width > 0); + assert!(tile_size.is_multiple_of(bank_width)); + assert!( + depth > 0, + "Matrix SRAM must contain at least one physical row" + ); + let banks = tile_size / bank_width; + assert!( + banks.is_power_of_two(), + "Matrix bank count must be a power of two" + ); + assert!(banks <= 64, "Matrix view skew has a 6-bit bank contract"); + let element_type = match ty { + MxDataType::Plain(data_type) => data_type, + MxDataType::Mx { .. } => { + panic!("Matrix SRAM bank words require a plain element type") + } + }; + assert!( + (bank_width as usize * element_type.size_in_bits() as usize).is_multiple_of(8), + "a Matrix bank word must contain whole bytes" + ); + + let full_tile_count = depth / tile_size as usize; + let words_per_row = tile_size / bank_width; + let rows_per_tile = tile_size * words_per_row.div_ceil(banks); + // MATRIX_SRAM_SIZE is measured in MLEN-wide rows. Views may use a + // proper subset of those rows even when the legacy square-tile API + // cannot fit one complete MLEN x MLEN tile. + let physical_rows = depth; + let bytes_per_word = + (bank_width as usize * element_type.size_in_bits() as usize).div_ceil(8); + let bank_rows = (0..banks) .map(|_| { - Mutex::new(Cell::Ready(QuantTensor::zeros( - (tile_size * tile_size) as usize, - ty, - ))) + (0..physical_rows) + .map(|_| Mutex::new(Cell::Ready(vec![0_u8; bytes_per_word]))) + .collect() }) .collect(); + Self { tile_size, - tiles, + depth_rows: depth, + full_tile_count, ty, + element_type, + banks, + bank_width, + rows_per_tile, + bank_rows, + packet_counters: MatrixPacketCounters::default(), } } @@ -35,96 +177,877 @@ impl MatrixSram { self.tile_size } + pub fn depth_rows(&self) -> usize { + self.depth_rows + } + pub fn ty(&self) -> MxDataType { self.ty } + pub fn banks(&self) -> u32 { + self.banks + } + + pub fn bank_width(&self) -> u32 { + self.bank_width + } + + pub fn element_bits(&self) -> u32 { + u32::from(self.element_type.size_in_bits()) + } + + /// Actual byte capacity of all physical bank words. pub fn size_in_bytes(&self) -> usize { - (self.tile_size * self.tile_size) as usize * self.tiles.len() + let bits = + self.depth_rows * self.tile_size as usize * self.element_type.size_in_bits() as usize; + bits.div_ceil(8) } - pub async fn read(&self, addr: u32) -> QuantTensor { - let idx = addr_to_cell(addr, self.tile_size * self.tile_size, self.tiles.len()); - tracing::trace!( - addr, - tile_idx = idx, - tile_size = self.tile_size, - "MRAM read" + pub fn default_layout(&self) -> MatrixLayout { + MatrixLayout { + rows: self.tile_size, + cols: self.tile_size, + tile_count: 1, + tile_pitch_rows: self.rows_per_tile, + alpha: 1, + tile_skew: 0, + } + } + + pub fn packet_counter_snapshot(&self) -> MatrixPacketCounterSnapshot { + MatrixPacketCounterSnapshot { + packets: self.packet_counters.packets.load(Ordering::Relaxed), + values: self.packet_counters.values.load(Ordering::Relaxed), + bank_words: self.packet_counters.bank_words.load(Ordering::Relaxed), + ideal_cycles: self.packet_counters.ideal_cycles.load(Ordering::Relaxed), + service_cycles: self.packet_counters.service_cycles.load(Ordering::Relaxed), + bank_stall_cycles: self + .packet_counters + .bank_stall_cycles + .load(Ordering::Relaxed), + } + } + + pub fn reset_packet_counters(&self) { + self.packet_counters.packets.store(0, Ordering::Relaxed); + self.packet_counters.values.store(0, Ordering::Relaxed); + self.packet_counters.bank_words.store(0, Ordering::Relaxed); + self.packet_counters + .ideal_cycles + .store(0, Ordering::Relaxed); + self.packet_counters + .service_cycles + .store(0, Ordering::Relaxed); + self.packet_counters + .bank_stall_cycles + .store(0, Ordering::Relaxed); + } + + pub fn physical_coord( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + row: u32, + col: u32, + ) -> MatrixPhysicalCoord { + self.validate_layout(addr, layout); + self.physical_coord_unchecked(addr, layout, tile, row, col) + } + + fn physical_coord_unchecked( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + row: u32, + col: u32, + ) -> MatrixPhysicalCoord { + assert!( + tile < layout.tile_count, + "Matrix-view tile index out of bounds" + ); + assert!(row < layout.rows, "Matrix-view row out of bounds"); + assert!(col < layout.cols, "Matrix-view column out of bounds"); + + let full_row_elements = self.banks * self.bank_width; + assert!( + addr.is_multiple_of(self.bank_width), + "Matrix-view base address {addr} is not aligned to a bank word" + ); + let base_bank_row = addr / full_row_elements; + let base_bank = (addr % full_row_elements) / self.bank_width; + let word = col / self.bank_width; + let words_per_row = layout.cols / self.bank_width; + let row_groups = words_per_row.div_ceil(self.banks); + let bank_row = + base_bank_row + tile * layout.tile_pitch_rows + row * row_groups + word / self.banks; + // The address already contains the allocation base, tile pitch, row, + // and wide-row word group. Using the tile-local `row` here discards + // that information. ISA-visible Matrix views always use alpha=1. + let bank = + (base_bank + layout.alpha * bank_row + layout.tile_skew * tile + word) % self.banks; + assert!( + (bank_row as usize) < self.bank_rows[bank as usize].len(), + "Matrix-view physical row {bank_row} exceeds bank capacity" ); - let mut guard = self.tiles[idx].lock().await; - guard.resolve().await.clone() + MatrixPhysicalCoord { + bank, + bank_row, + lane: col % self.bank_width, + } + } + + pub async fn read(&self, addr: u32) -> QuantTensor { + self.read_layout_tile(addr, self.default_layout(), 0).await } pub async fn write(&self, addr: u32, tensor: QuantTensor) { - let idx = addr_to_cell(addr, self.tile_size * self.tile_size, self.tiles.len()); - assert!(tensor.data_type() == self.ty); - *self.tiles[idx].lock().await = Cell::Ready(tensor); + assert_eq!( + tensor.data_type(), + self.ty, + "legacy Matrix write must match the SRAM data type" + ); + self.write_layout_tile(addr, self.default_layout(), 0, tensor) + .await; } - pub async fn write_delayed(&self, addr: u32, tensor: Receiver) { - // PRE-EXISTING ODDITY: divides by `tile_size` rather than the - // `tile_size * tile_size` used by every other method. Likely a bug, - // but preserved verbatim from the original implementation pending - // dedicated investigation. - let idx = addr_to_cell(addr, self.tile_size, self.tiles.len()); - *self.tiles[idx].lock().await = Cell::Pending(tensor); - } - - /// Park `cells` consecutive tiles starting at element address `addr` as - /// [`Cell::Pending`], each on its own channel; returns the senders in cell - /// order. Non-blocking counterpart of [`Self::continous_write_delayed`]: - /// the caller spawns [`Self::fill_pending`] to feed the channels while - /// readers of the parked cells block until their chunk arrives. + /// Read one logical tile through a configured placement and restore lanes. + pub async fn read_layout_tile( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + ) -> QuantTensor { + self.validate_layout(addr, layout); + let mut logical = vec![0_f32; (layout.rows * layout.cols) as usize]; + let words_per_row = layout.cols / self.bank_width; + for row in 0..layout.rows { + for word in 0..words_per_row { + let col = word * self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + let values = self.word_bytes_to_values(&bytes); + let start = (row * layout.cols + col) as usize; + logical[start..start + self.bank_width as usize].copy_from_slice(&values); + } + } + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty) + } + + /// Read one logical row or column from the same physical cells. /// - /// Uses the same `tile_size * tile_size` cell divisor as `read`/`write` - /// (not `write_delayed`'s flagged `tile_size` divisor). - pub async fn mark_pending_tiles( + /// A bank contributes at most one word per cycle. Column reads select one + /// lane from each returned bank word and restore logical row order; there + /// is no transposed copy or hidden transpose buffer. + pub async fn read_layout_line( &self, addr: u32, - cells: u32, - ) -> Vec> { - let start_idx = addr_to_cell(addr, self.tile_size * self.tile_size, self.tiles.len()); - let mut senders = Vec::with_capacity(cells as usize); - for i in 0..cells as usize { - let (tx, rx) = tokio::sync::oneshot::channel(); - *self.tiles[start_idx + i].lock().await = Cell::Pending(rx); - senders.push(tx); + layout: MatrixLayout, + tile: u32, + index: u32, + axis: MatrixAccessAxis, + ) -> (QuantTensor, MatrixPacketService) { + self.validate_layout(addr, layout); + let positions = match axis { + MatrixAccessAxis::Row => { + assert!(index < layout.rows, "Matrix-view row out of bounds"); + (0..layout.cols) + .step_by(self.bank_width as usize) + .map(|col| (index, col, true)) + .collect::>() + } + MatrixAccessAxis::Column => { + assert!(index < layout.cols, "Matrix-view column out of bounds"); + (0..layout.rows) + .map(|row| (row, index, false)) + .collect::>() + } + }; + + let mut logical = Vec::with_capacity(match axis { + MatrixAccessAxis::Row => layout.cols as usize, + MatrixAccessAxis::Column => layout.rows as usize, + }); + let mut per_bank = vec![0_u64; self.banks as usize]; + for (row, col, whole_word) in positions { + let word_col = col - col % self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, word_col); + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + let values = self.word_bytes_to_values(&bytes); + if whole_word { + logical.extend(values); + } else { + logical.push(values[(col % self.bank_width) as usize]); + } + per_bank[coord.bank as usize] += 1; } - senders + let service = self.packet_service(&per_bank, logical.len() as u64); + self.record_packet(service); + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + service, + ) } - /// Completer half of [`Self::mark_pending_tiles`]: awaits the DMA tensor, - /// splits it into tile-sized chunks exactly like - /// [`Self::continous_write_delayed`], and feeds each parked cell's channel. - /// On DMA failure the senders drop and the cells stay pending, so a later - /// reader fails loudly instead of seeing stale data. - pub async fn fill_pending( + /// Read several logical rows or columns in one bank-service packet. + /// + /// For column groups, lanes that share a physical bank word are fetched + /// once and then restored to their logical columns. The service record is + /// derived from those exact physical words, so diagonal-read correctness + /// and bank-conflict timing cannot diverge. + pub async fn read_layout_lines( &self, - senders: Vec>, - tensor: Receiver, - ) { - match tensor.await { - Ok(tensor) => { - let chunk_size = (self.tile_size * self.tile_size) as i64; - let total = tensor.as_tensor().size()[0]; - for (i, sender) in senders.into_iter().enumerate() { - let start = (i as i64) * chunk_size; - if start >= total { - break; + addr: u32, + layout: MatrixLayout, + tile: u32, + first: u32, + count: u32, + axis: MatrixAccessAxis, + ) -> (QuantTensor, MatrixPacketService) { + self.validate_layout(addr, layout); + assert!(count > 0, "Matrix line packet must be non-empty"); + let limit = match axis { + MatrixAccessAxis::Row => layout.rows, + MatrixAccessAxis::Column => layout.cols, + }; + assert!( + first + count <= limit, + "Matrix line packet exceeds its view" + ); + + let line_len = match axis { + MatrixAccessAxis::Row => layout.cols, + MatrixAccessAxis::Column => layout.rows, + }; + let mut logical = vec![0_f32; (count * line_len) as usize]; + let mut per_bank = vec![0_u64; self.banks as usize]; + let mut words: HashMap<(u32, u32), Vec> = HashMap::new(); + + for line_offset in 0..count { + let line = first + line_offset; + match axis { + MatrixAccessAxis::Row => { + for col in (0..layout.cols).step_by(self.bank_width as usize) { + let coord = self.physical_coord_unchecked(addr, layout, tile, line, col); + let key = (coord.bank, coord.bank_row); + if let Entry::Vacant(entry) = words.entry(key) { + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize] + [coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + entry.insert(self.word_bytes_to_values(&bytes)); + per_bank[coord.bank as usize] += 1; + } + let values = &words[&key]; + let start = (line_offset * line_len + col) as usize; + logical[start..start + self.bank_width as usize].copy_from_slice(values); + } + } + MatrixAccessAxis::Column => { + for row in 0..layout.rows { + let word_col = line - line % self.bank_width; + let coord = + self.physical_coord_unchecked(addr, layout, tile, row, word_col); + let key = (coord.bank, coord.bank_row); + if let Entry::Vacant(entry) = words.entry(key) { + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize] + [coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + entry.insert(self.word_bytes_to_values(&bytes)); + per_bank[coord.bank as usize] += 1; + } + logical[(line_offset * line_len + row) as usize] = + words[&key][(line % self.bank_width) as usize]; } - let end = ((i as i64 + 1) * chunk_size).min(total); - let chunk = tensor - .as_tensor() - .narrow(0, start, end - start) - .shallow_clone(); - let chunk_qt = QuantTensor::quantize(chunk, self.ty); - let _ = sender.send(chunk_qt); } } - Err(_) => { - tracing::error!("delayed matrix fill skipped: DMA sender dropped"); + } + + let service = self.packet_service(&per_bank, logical.len() as u64); + self.record_packet(service); + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + service, + ) + } + + /// Read an explicitly ordered packet of logical rows across one or more + /// tiles. `lines` contains `(tile, row)` pairs and the returned tensor is + /// the concatenation of those rows in exactly that order. + /// + /// L-Tile recurrence packets use row-major tile order: for a fixed state + /// index they request the same logical row from several head tiles. Bank + /// service is computed from the physical words that supplied those values, + /// including de-duplication for an explicitly broadcast source row. + pub async fn read_layout_indexed_rows( + &self, + addr: u32, + layout: MatrixLayout, + lines: &[(u32, u32)], + ) -> (QuantTensor, MatrixPacketService) { + self.validate_layout(addr, layout); + assert!( + !lines.is_empty(), + "Matrix indexed-row packet must be non-empty" + ); + + let mut logical = vec![0_f32; lines.len() * layout.cols as usize]; + let mut per_bank = vec![0_u64; self.banks as usize]; + let mut words: HashMap<(u32, u32), Vec> = HashMap::new(); + + for (line_index, &(tile, row)) in lines.iter().enumerate() { + assert!( + tile < layout.tile_count, + "Matrix indexed-row tile out of bounds" + ); + assert!(row < layout.rows, "Matrix indexed-row row out of bounds"); + for col in (0..layout.cols).step_by(self.bank_width as usize) { + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let key = (coord.bank, coord.bank_row); + if let Entry::Vacant(entry) = words.entry(key) { + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize] + [coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + entry.insert(self.word_bytes_to_values(&bytes)); + per_bank[coord.bank as usize] += 1; + } + let start = line_index * layout.cols as usize + col as usize; + logical[start..start + self.bank_width as usize].copy_from_slice(&words[&key]); } } + + let service = self.packet_service(&per_bank, logical.len() as u64); + self.record_packet(service); + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + service, + ) + } + + /// Read an explicitly ordered packet of logical columns across tiles. + /// + /// This is the transpose counterpart of `read_layout_indexed_rows`. It + /// reads the same physical cells and returns `[requested_line][logical_row]` + /// order. A head-major `[head][key]` field can therefore feed a key-major + /// recurrence packet without a copied transpose. + pub async fn read_layout_indexed_columns( + &self, + addr: u32, + layout: MatrixLayout, + lines: &[(u32, u32)], + ) -> (QuantTensor, MatrixPacketService) { + self.validate_layout(addr, layout); + assert!( + !lines.is_empty(), + "Matrix indexed-column packet must be non-empty" + ); + + let mut logical = vec![0_f32; lines.len() * layout.rows as usize]; + let mut per_bank = vec![0_u64; self.banks as usize]; + let mut words: HashMap<(u32, u32), Vec> = HashMap::new(); + + for (line_index, &(tile, col)) in lines.iter().enumerate() { + assert!( + tile < layout.tile_count, + "Matrix indexed-column tile out of bounds" + ); + assert!( + col < layout.cols, + "Matrix indexed-column column out of bounds" + ); + let word_col = col - col % self.bank_width; + for row in 0..layout.rows { + let coord = self.physical_coord_unchecked(addr, layout, tile, row, word_col); + let key = (coord.bank, coord.bank_row); + if let Entry::Vacant(entry) = words.entry(key) { + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize] + [coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + entry.insert(self.word_bytes_to_values(&bytes)); + per_bank[coord.bank as usize] += 1; + } + logical[line_index * layout.rows as usize + row as usize] = + words[&key][(col % self.bank_width) as usize]; + } + } + + let service = self.packet_service(&per_bank, logical.len() as u64); + self.record_packet(service); + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + service, + ) + } + + /// Write several logical rows from a lane-restored packet. + pub async fn write_layout_rows( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + first_row: u32, + row_count: u32, + tensor: QuantTensor, + ) -> MatrixPacketService { + self.validate_layout(addr, layout); + assert!(row_count > 0 && first_row + row_count <= layout.rows); + let values = tensor_to_f32_vec(tensor.as_tensor()); + assert_eq!(values.len(), (row_count * layout.cols) as usize); + let mut per_bank = vec![0_u64; self.banks as usize]; + for row_offset in 0..row_count { + let row = first_row + row_offset; + for col in (0..layout.cols).step_by(self.bank_width as usize) { + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let start = (row_offset * layout.cols + col) as usize; + let bytes = + self.values_to_word_bytes(&values[start..start + self.bank_width as usize]); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Ready(bytes); + per_bank[coord.bank as usize] += 1; + } + } + let service = self.packet_service(&per_bank, values.len() as u64); + self.record_packet(service); + service + } + + /// Write an explicitly ordered packet of logical rows across head tiles. + /// Distinct destination words are required to map to distinct physical + /// cells; otherwise a purported layout silently aliases recurrent state. + pub async fn write_layout_indexed_rows( + &self, + addr: u32, + layout: MatrixLayout, + lines: &[(u32, u32)], + tensor: QuantTensor, + ) -> MatrixPacketService { + self.validate_layout(addr, layout); + assert!( + !lines.is_empty(), + "Matrix indexed-row write must be non-empty" + ); + let values = tensor_to_f32_vec(tensor.as_tensor()); + assert_eq!(values.len(), lines.len() * layout.cols as usize); + let mut per_bank = vec![0_u64; self.banks as usize]; + let mut destinations = HashMap::new(); + + for (line_index, &(tile, row)) in lines.iter().enumerate() { + assert!( + tile < layout.tile_count, + "Matrix indexed-row tile out of bounds" + ); + assert!(row < layout.rows, "Matrix indexed-row row out of bounds"); + for col in (0..layout.cols).step_by(self.bank_width as usize) { + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let key = (coord.bank, coord.bank_row); + assert!( + destinations.insert(key, (tile, row, col)).is_none(), + "Matrix view aliases two logical destination words at bank {} row {}", + coord.bank, + coord.bank_row + ); + let start = line_index * layout.cols as usize + col as usize; + let bytes = + self.values_to_word_bytes(&values[start..start + self.bank_width as usize]); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Ready(bytes); + per_bank[coord.bank as usize] += 1; + } + } + + let service = self.packet_service(&per_bank, values.len() as u64); + self.record_packet(service); + service + } + + /// Reconstruct a complete tile while accounting for row- or column-wise + /// service. The returned tensor is always logical row-major data. + pub async fn read_layout_tile_axis( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + axis: MatrixAccessAxis, + ) -> (QuantTensor, MatrixPacketService) { + let mut logical = vec![0_f32; (layout.rows * layout.cols) as usize]; + let mut total = MatrixPacketService::default(); + let lines = match axis { + MatrixAccessAxis::Row => layout.rows, + MatrixAccessAxis::Column => layout.cols, + }; + for index in 0..lines { + let (line, service) = self.read_layout_line(addr, layout, tile, index, axis).await; + let values = tensor_to_f32_vec(line.as_tensor()); + match axis { + MatrixAccessAxis::Row => { + let start = (index * layout.cols) as usize; + logical[start..start + layout.cols as usize].copy_from_slice(&values); + } + MatrixAccessAxis::Column => { + for (row, value) in values.into_iter().enumerate() { + logical[row * layout.cols as usize + index as usize] = value; + } + } + } + total.values += service.values; + total.bank_words += service.bank_words; + total.ideal_cycles += service.ideal_cycles; + total.service_cycles += service.service_cycles; + total.bank_stall_cycles += service.bank_stall_cycles; + total.worst_bank_words = total.worst_bank_words.max(service.worst_bank_words); + } + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + total, + ) + } + + /// Write one logical tile through the affine mapper into physical banks. + pub async fn write_layout_tile( + &self, + addr: u32, + layout: MatrixLayout, + tile: u32, + tensor: QuantTensor, + ) { + self.validate_layout(addr, layout); + let mut logical = tensor_to_f32_vec(tensor.as_tensor()); + let expected = (layout.rows * layout.cols) as usize; + assert!( + logical.len() <= expected, + "Matrix tile is larger than its view" + ); + logical.resize(expected, 0.0); + let words_per_row = layout.cols / self.bank_width; + for row in 0..layout.rows { + for word in 0..words_per_row { + let col = word * self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let start = (row * layout.cols + col) as usize; + let bytes = + self.values_to_word_bytes(&logical[start..start + self.bank_width as usize]); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Ready(bytes); + } + } + } + + /// Read every tile in one logical packet and restore tile/row/column order. + /// + /// One bank word per bank may be served each cycle. The returned service + /// record is calculated from the same physical coordinates that supplied + /// the returned values, so timing and correctness cannot drift apart. + pub async fn read_layout_packet( + &self, + addr: u32, + layout: MatrixLayout, + ) -> (QuantTensor, MatrixPacketService) { + let (packet, per_bank) = self.read_layout_packet_raw(addr, layout).await; + let service = self.packet_service(&per_bank, packet.as_tensor().numel() as u64); + self.record_packet(service); + (packet, service) + } + + /// Read several logical operands in one Vector issue slot. + /// + /// Their bank loads are combined before service time is calculated. This is + /// what distinguishes a genuinely same-cycle cross-field access from a list + /// of independent one-packet microbenchmarks. + pub async fn read_layout_packets( + &self, + requests: &[(u32, MatrixLayout)], + ) -> (Vec, MatrixPacketService) { + assert!( + !requests.is_empty(), + "a Matrix packet group cannot be empty" + ); + let mut packets = Vec::with_capacity(requests.len()); + let mut per_bank = vec![0_u64; self.banks as usize]; + let mut values = 0_u64; + for &(addr, layout) in requests { + let (packet, loads) = self.read_layout_packet_raw(addr, layout).await; + values += packet.as_tensor().numel() as u64; + for (total, load) in per_bank.iter_mut().zip(loads) { + *total += load; + } + packets.push(packet); + } + let service = self.packet_service(&per_bank, values); + self.record_packet_count(service, requests.len() as u64); + (packets, service) + } + + /// Scatter one tile-major packet through the same affine map used by read. + pub async fn write_layout_packet( + &self, + addr: u32, + layout: MatrixLayout, + tensor: QuantTensor, + ) -> MatrixPacketService { + self.validate_layout(addr, layout); + let expected = (layout.tile_count * layout.rows * layout.cols) as usize; + let values = tensor_to_f32_vec(tensor.as_tensor()); + assert_eq!( + values.len(), + expected, + "Matrix packet contains {} values, view requires {expected}", + values.len() + ); + let words_per_row = layout.cols / self.bank_width; + let mut per_bank = vec![0_u64; self.banks as usize]; + for tile in 0..layout.tile_count { + for row in 0..layout.rows { + for word in 0..words_per_row { + let col = word * self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let start = ((tile * layout.rows + row) * layout.cols + col) as usize; + let bytes = + self.values_to_word_bytes(&values[start..start + self.bank_width as usize]); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Ready(bytes); + per_bank[coord.bank as usize] += 1; + } + } + } + let bank_words = per_bank.iter().sum::(); + let service_cycles = per_bank.iter().copied().max().unwrap_or(0); + let ideal_cycles = bank_words.div_ceil(u64::from(self.banks)); + let service = MatrixPacketService { + values: values.len() as u64, + bank_words, + ideal_cycles, + service_cycles, + bank_stall_cycles: service_cycles.saturating_sub(ideal_cycles), + worst_bank_words: service_cycles, + }; + self.record_packet(service); + service + } + + /// Write one dense microtile into a logical view at `logical_offset`. + /// + /// This is the direct Matrix-accumulator writeback path. The offset is in + /// logical elements, so physical placement remains entirely a property of + /// the configured view. One accumulator row is exactly one physical bank + /// word; keeping that invariant explicit avoids hiding a crossbar or a + /// read-modify-write in the timing model. + pub async fn write_layout_microtile( + &self, + addr: u32, + layout: MatrixLayout, + logical_offset: u32, + tensor: QuantTensor, + micro_rows: u32, + micro_cols: u32, + ) -> MatrixPacketService { + self.validate_layout(addr, layout); + assert_eq!( + micro_cols, self.bank_width, + "Matrix accumulator row must equal one Matrix bank word" + ); + let values = tensor_to_f32_vec(tensor.as_tensor()); + assert_eq!(values.len(), (micro_rows * micro_cols) as usize); + let tile_values = layout.rows * layout.cols; + let total_values = tile_values * layout.tile_count; + assert!(logical_offset < total_values); + + let start_in_tile = logical_offset % tile_values; + let start_col = start_in_tile % layout.cols; + assert!(start_col.is_multiple_of(self.bank_width)); + assert!(start_col + micro_cols <= layout.cols); + + let mut per_bank = vec![0_u64; self.banks as usize]; + for micro_row in 0..micro_rows { + let flat = logical_offset + micro_row * layout.cols; + assert!(flat + micro_cols <= total_values); + let tile = flat / tile_values; + let within = flat % tile_values; + let row = within / layout.cols; + let col = within % layout.cols; + assert!(row < layout.rows); + assert_eq!(col, start_col, "microtile may not wrap a logical row"); + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + assert_eq!(coord.lane, 0); + let value_start = (micro_row * micro_cols) as usize; + let bytes = + self.values_to_word_bytes(&values[value_start..value_start + micro_cols as usize]); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Ready(bytes); + per_bank[coord.bank as usize] += 1; + } + + let bank_words = per_bank.iter().sum::(); + let service_cycles = per_bank.iter().copied().max().unwrap_or(0); + let ideal_cycles = bank_words.div_ceil(u64::from(self.banks)); + let service = MatrixPacketService { + values: values.len() as u64, + bank_words, + ideal_cycles, + service_cycles, + bank_stall_cycles: service_cycles.saturating_sub(ideal_cycles), + worst_bank_words: service_cycles, + }; + self.record_packet(service); + service + } + + /// Preserve the historical odd address divisor of this legacy API. + pub async fn write_delayed(&self, addr: u32, tensor: Receiver) { + assert!( + self.full_tile_count > 0, + "legacy Matrix tile API requires at least MLEN physical rows" + ); + let index = addr_to_cell(addr, self.tile_size, self.full_tile_count); + let tensor = tensor.await.expect("delayed Matrix write sender dropped"); + self.write(index as u32 * self.tile_size * self.tile_size, tensor) + .await; + } + + /// Mark every physical bank word in `cells` default-layout tiles pending. + pub async fn mark_pending_tiles(&self, addr: u32, cells: u32) -> PendingMatrixTiles { + assert!( + self.full_tile_count > 0, + "legacy Matrix tile API requires at least MLEN physical rows" + ); + let start_tile = addr_to_cell(addr, self.tile_size * self.tile_size, self.full_tile_count); + let count = (cells as usize).min(self.full_tile_count.saturating_sub(start_tile)); + let layout = MatrixLayout { + tile_count: count as u32, + ..self.default_layout() + }; + let mut pending = Vec::with_capacity(count * self.tile_size as usize * self.banks as usize); + for tile in 0..count { + for row in 0..self.tile_size as usize { + for word in 0..self.banks as usize { + let col = word as u32 * self.bank_width; + let coord = + self.physical_coord_unchecked(addr, layout, tile as u32, row as u32, col); + let (sender, receiver) = tokio::sync::oneshot::channel(); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Pending(receiver); + pending.push(PendingWord { + sender, + source_offset: tile * self.tile_size as usize * self.tile_size as usize + + row * self.tile_size as usize + + word * self.bank_width as usize, + }); + } + } + } + PendingMatrixTiles { words: pending } + } + + /// Park every physical bank word selected by an explicit Matrix view. + /// + /// The completion handle preserves each word's logical packet offset, so + /// the DMA result is scattered through the exact same affine map later + /// consumed by `read_layout_packet`/`L_TILE_EXEC`. + pub async fn mark_pending_layout_packet( + &self, + addr: u32, + layout: MatrixLayout, + ) -> (PendingMatrixTiles, MatrixPacketService) { + self.validate_layout(addr, layout); + let words_per_row = layout.cols / self.bank_width; + let mut pending = + Vec::with_capacity((layout.tile_count * layout.rows * words_per_row) as usize); + let mut destinations = HashMap::new(); + let mut per_bank = vec![0_u64; self.banks as usize]; + for tile in 0..layout.tile_count { + for row in 0..layout.rows { + for word in 0..words_per_row { + let col = word * self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + assert!( + destinations + .insert((coord.bank, coord.bank_row), (tile, row, word)) + .is_none(), + "Matrix view aliases pending DMA destinations" + ); + let (sender, receiver) = tokio::sync::oneshot::channel(); + *self.bank_rows[coord.bank as usize][coord.bank_row as usize] + .lock() + .await = Cell::Pending(receiver); + pending.push(PendingWord { + source_offset: ((tile * layout.rows + row) * layout.cols + col) as usize, + sender, + }); + per_bank[coord.bank as usize] += 1; + } + } + } + let values = u64::from(layout.tile_count) * u64::from(layout.rows) * u64::from(layout.cols); + let service = self.packet_service(&per_bank, values); + self.record_packet(service); + (PendingMatrixTiles { words: pending }, service) + } + + pub async fn fill_pending(&self, pending: PendingMatrixTiles, tensor: Receiver) { + let tensor = tensor + .await + .unwrap_or_else(|error| panic!("delayed Matrix fill sender dropped: {error}")); + let values = tensor_to_f32_vec(tensor.as_tensor()); + for word in pending.words { + let source = word.source_offset; + let mut padded = vec![0_f32; self.bank_width as usize]; + if source < values.len() { + let end = (source + self.bank_width as usize).min(values.len()); + padded[..end - source].copy_from_slice(&values[source..end]); + } + let quantized = QuantTensor::quantize(tensor_from_f32_slice(&padded), self.ty); + let _ = word.sender.send(quantized); + } } pub async fn continous_write_delayed( @@ -133,61 +1056,234 @@ impl MatrixSram { write_amount: u32, tensor: Receiver, ) { - let start_idx = addr_to_cell(addr, self.tile_size * self.tile_size, self.tiles.len()); - // Await the tensor from the channel (blocks until data arrives) - if let Ok(tensor) = tensor.await { - let dims = tensor.as_tensor().size(); - let chunk_size = (self.tile_size * self.tile_size) as i64; - let total = dims[0]; - - // Split the tensor into chunks of self.tile_size and store each in self.tiles. - for i in 0..write_amount.min((total as u32).div_ceil(self.tile_size * self.tile_size)) { - let start = (i as i64) * chunk_size; - let end = ((i as i64 + 1) * chunk_size).min(total); - let chunk = tensor - .as_tensor() - .narrow(0, start, end - start) - .shallow_clone(); - let chunk_qt = QuantTensor::quantize(chunk, self.ty); - *self.tiles[start_idx + i as usize].lock().await = Cell::Ready(chunk_qt); - } - } else { - // The DMA producer dropped its sender: the prefetch was - // cancelled/failed, so these cells keep their previous contents. - tracing::error!( - addr, - write_amount, - "delayed matrix write skipped: DMA sender dropped" - ); + let tensor = tensor + .await + .unwrap_or_else(|error| panic!("delayed Matrix write sender dropped: {error}")); + let values = tensor_to_f32_vec(tensor.as_tensor()); + let tile_elements = (self.tile_size * self.tile_size) as usize; + let count = (write_amount as usize) + .min(values.len().div_ceil(tile_elements)) + .min(self.full_tile_count); + for tile in 0..count { + let start = tile * tile_elements; + let end = (start + tile_elements).min(values.len()); + let mut padded = vec![0_f32; tile_elements]; + padded[..end - start].copy_from_slice(&values[start..end]); + let quantized = QuantTensor::quantize(tensor_from_f32_slice(&padded), self.ty); + self.write( + addr + tile as u32 * self.tile_size * self.tile_size, + quantized, + ) + .await; } } pub async fn as_bytes(&self) -> Vec { - let element_ty = self.ty.element_type(); - let mut result = Vec::new(); - - for tile_mutex in &self.tiles { - let mut guard = tile_mutex.lock().await; - let tensor = guard.resolve().await; - let tensor_data = tensor.as_tensor(); - let f32_vec = tensor_to_f32_vec(tensor_data); - let len = f32_vec.len(); - // Calculate bytes needed for THIS tile's actual size - let total_bits = len * element_ty.size_in_bits() as usize; - let bytes_needed = total_bits.div_ceil(8); - let mut tile_bytes = vec![0u8; bytes_needed]; - element_ty.bytes_from_f32(&f32_vec, &mut tile_bytes); - result.extend_from_slice(&tile_bytes); + let mut result = Vec::with_capacity(self.size_in_bytes()); + // Export the complete physical capacity in default logical row order, + // including a final partial-square region. This keeps dump size equal + // to MATRIX_SRAM_SIZE * MLEN rather than silently dropping spare rows. + for row in 0..self.depth_rows as u32 { + for word in 0..self.banks { + let bank = (row + word) % self.banks; + let bytes = { + let mut guard = self.bank_rows[bank as usize][row as usize].lock().await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + result.extend_from_slice(&bytes); + } } - + debug_assert_eq!(result.len(), self.size_in_bytes()); result } + + fn validate_layout(&self, addr: u32, layout: MatrixLayout) { + assert!(layout.rows > 0 && layout.cols > 0 && layout.tile_count > 0); + assert!(layout.cols.is_multiple_of(self.bank_width)); + assert!(layout.alpha < self.banks); + let words_per_row = layout.cols / self.bank_width; + let row_groups = words_per_row.div_ceil(self.banks); + let full_row_elements = self.banks * self.bank_width; + assert!(addr.is_multiple_of(self.bank_width)); + let base_bank_row = addr / full_row_elements; + let base_bank = (addr % full_row_elements) / self.bank_width; + let mut occupied = HashMap::new(); + let mut final_row = base_bank_row; + for tile in 0..layout.tile_count { + for row in 0..layout.rows { + for word in 0..words_per_row { + let bank_row = base_bank_row + + tile * layout.tile_pitch_rows + + row * row_groups + + word / self.banks; + let bank = + (base_bank + layout.alpha * bank_row + layout.tile_skew * tile + word) + % self.banks; + let previous = occupied.insert((bank, bank_row), (tile, row, word)); + assert!( + previous.is_none(), + "Matrix view aliases logical words {:?} and {:?} at bank {bank}, row {bank_row}", + previous.unwrap_or_default(), + (tile, row, word), + ); + final_row = final_row.max(bank_row + 1); + } + } + } + assert!( + final_row <= self.bank_rows[0].len() as u32, + "Matrix view exceeds SRAM capacity" + ); + } + + /// Check that several compiler-managed views can coexist in this SRAM. + /// + /// A descriptor proves only that a tensor does not alias itself. A static + /// allocator must additionally prove that two live tensors never claim the + /// same physical bank word. This is deliberately not a cache check: there + /// are no tags, misses or replacement decisions, only explicit addresses. + pub fn validate_disjoint_layouts( + &self, + views: &[(&str, u32, MatrixLayout)], + ) -> Result<(), String> { + let mut occupied: HashMap<(u32, u32), (&str, u32, u32, u32)> = HashMap::new(); + for &(name, addr, layout) in views { + self.validate_layout(addr, layout); + let words_per_row = layout.cols / self.bank_width; + for tile in 0..layout.tile_count { + for row in 0..layout.rows { + for word in 0..words_per_row { + let coord = self.physical_coord_unchecked( + addr, + layout, + tile, + row, + word * self.bank_width, + ); + let key = (coord.bank, coord.bank_row); + if let Some(previous) = occupied.insert(key, (name, tile, row, word)) { + return Err(format!( + "Matrix views alias physical bank word bank={}, row={}: {:?} and {:?}", + coord.bank, + coord.bank_row, + previous, + (name, tile, row, word), + )); + } + } + } + } + } + Ok(()) + } + + async fn read_layout_packet_raw( + &self, + addr: u32, + layout: MatrixLayout, + ) -> (QuantTensor, Vec) { + self.validate_layout(addr, layout); + let mut logical = + Vec::with_capacity((layout.tile_count * layout.rows * layout.cols) as usize); + let words_per_row = layout.cols / self.bank_width; + let mut per_bank = vec![0_u64; self.banks as usize]; + for tile in 0..layout.tile_count { + for row in 0..layout.rows { + for word in 0..words_per_row { + let col = word * self.bank_width; + let coord = self.physical_coord_unchecked(addr, layout, tile, row, col); + let bytes = { + let mut guard = self.bank_rows[coord.bank as usize] + [coord.bank_row as usize] + .lock() + .await; + guard + .resolve_with(|tensor| self.tensor_to_word_bytes(&tensor)) + .await + .clone() + }; + logical.extend(self.word_bytes_to_values(&bytes)); + per_bank[coord.bank as usize] += 1; + } + } + } + ( + QuantTensor::quantize(tensor_from_f32_slice(&logical), self.ty), + per_bank, + ) + } + + fn packet_service(&self, per_bank: &[u64], values: u64) -> MatrixPacketService { + let bank_words = per_bank.iter().sum::(); + let service_cycles = per_bank.iter().copied().max().unwrap_or(0); + let ideal_cycles = bank_words.div_ceil(u64::from(self.banks)); + MatrixPacketService { + values, + bank_words, + ideal_cycles, + service_cycles, + bank_stall_cycles: service_cycles.saturating_sub(ideal_cycles), + worst_bank_words: service_cycles, + } + } + + fn record_packet(&self, service: MatrixPacketService) { + self.record_packet_count(service, 1); + } + + fn record_packet_count(&self, service: MatrixPacketService, packets: u64) { + self.packet_counters + .packets + .fetch_add(packets, Ordering::Relaxed); + self.packet_counters + .values + .fetch_add(service.values, Ordering::Relaxed); + self.packet_counters + .bank_words + .fetch_add(service.bank_words, Ordering::Relaxed); + self.packet_counters + .ideal_cycles + .fetch_add(service.ideal_cycles, Ordering::Relaxed); + self.packet_counters + .service_cycles + .fetch_add(service.service_cycles, Ordering::Relaxed); + self.packet_counters + .bank_stall_cycles + .fetch_add(service.bank_stall_cycles, Ordering::Relaxed); + } + + fn bytes_per_word(&self) -> usize { + (self.bank_width as usize * self.element_type.size_in_bits() as usize).div_ceil(8) + } + + fn tensor_to_word_bytes(&self, tensor: &QuantTensor) -> Vec { + let mut values = tensor_to_f32_vec(tensor.as_tensor()); + values.resize(self.bank_width as usize, 0.0); + values.truncate(self.bank_width as usize); + self.values_to_word_bytes(&values) + } + + fn values_to_word_bytes(&self, values: &[f32]) -> Vec { + let mut bytes = vec![0_u8; self.bytes_per_word()]; + self.element_type.bytes_from_f32(values, &mut bytes); + bytes + } + + fn word_bytes_to_values(&self, bytes: &[u8]) -> Vec { + let mut values = vec![0_f32; self.bank_width as usize]; + self.element_type + .convert_bytes_to_f32_vec(bytes, &mut values); + values + } } #[cfg(test)] mod tests { use super::*; - use quantize::{DataType, FpType}; + use quantize::{FpType, MxDataType}; use tch::Tensor; use tokio::sync::oneshot; @@ -195,6 +1291,10 @@ mod tests { MxDataType::Plain(DataType::Fp(FpType::F32)) } + fn bf16_plain() -> MxDataType { + MxDataType::Plain(DataType::Fp(FpType::BF16)) + } + fn tile(ty: MxDataType, vals: &[f32]) -> QuantTensor { QuantTensor::new_assuming_quantized(Tensor::from_slice(vals), ty).unwrap() } @@ -203,25 +1303,72 @@ mod tests { fn test_matrix_new_dimensions() { let m = MatrixSram::new(2, 8, f32_plain()); assert_eq!(m.tile_size(), 2); - // cells = depth / tile_size = 4; size = tile_size^2 * cells = 4 * 4. - assert_eq!(m.size_in_bytes(), 16); + assert_eq!(m.banks(), 2); + assert_eq!(m.bank_width(), 1); + assert_eq!(m.size_in_bytes(), 64); + } + + #[test] + fn default_constructor_supports_paper_mlen() { + let m = MatrixSram::new(2048, 256, bf16_plain()); + assert_eq!(m.banks(), 64); + assert_eq!(m.bank_width(), 32); + assert_eq!(m.size_in_bytes(), 1024 * 1024); } #[tokio::test] async fn test_matrix_write_read_roundtrip() { let ty = f32_plain(); - let m = MatrixSram::new(2, 8, ty); // tile_size 2 -> 4 elements per tile + let m = MatrixSram::new(2, 8, ty); let qt = tile(ty, &[1.0, 2.0, 3.0, 4.0]); - m.write(4, qt.clone()).await; // addr 4 -> cell 1 (4 / tile_size^2) + m.write(4, qt.clone()).await; let got = m.read(4).await; assert!(got.as_tensor().equal(qt.as_tensor())); + assert_eq!( + m.packet_counter_snapshot(), + MatrixPacketCounterSnapshot::default() + ); + } + + #[tokio::test] + #[should_panic(expected = "legacy Matrix write must match the SRAM data type")] + async fn legacy_write_rejects_a_mismatched_element_type() { + let m = MatrixSram::new(2, 8, bf16_plain()); + m.write(0, tile(f32_plain(), &[1.0, 2.0, 3.0, 4.0])).await; + } + + #[tokio::test] + async fn paper_depth_is_mlen_wide_rows_and_supports_compact_views() { + const MLEN: u32 = 2048; + const DEPTH_ROWS: usize = 256; + const BLEN: u32 = 32; + let ty = bf16_plain(); + let m = MatrixSram::with_banks(MLEN, DEPTH_ROWS, BLEN, ty); + assert_eq!(m.depth_rows(), DEPTH_ROWS); + assert_eq!(m.size_in_bytes(), 1024 * 1024); + + let view = MatrixLayout { + rows: 1, + cols: 64, + tile_count: 32, + tile_pitch_rows: 1, + alpha: 2, + tile_skew: 0, + }; + let values = (0..MLEN) + .map(|index| ((index % 127) as f32 - 63.0) / 16.0) + .collect::>(); + let input = QuantTensor::quantize(Tensor::from_slice(&values), ty); + let expected = tensor_to_f32_vec(input.as_tensor()); + let write = m.write_layout_packet(0, view, input).await; + let (output, read) = m.read_layout_packet(0, view).await; + assert_eq!(tensor_to_f32_vec(output.as_tensor()), expected); + assert_eq!((write.service_cycles, read.service_cycles), (1, 1)); + assert_eq!(read.bank_stall_cycles, 0); } #[tokio::test] - async fn test_matrix_write_delayed_uses_tile_size_divisor() { - // write_delayed divides the address by tile_size (2), while read/write - // divide by tile_size^2 (4). So write_delayed(2) and read(4) address the - // same cell (index 1). This pins that pre-existing oddity. + async fn test_matrix_write_delayed_preserves_legacy_address_divisor() { let ty = f32_plain(); let m = MatrixSram::new(2, 8, ty); let qt = tile(ty, &[5.0, 6.0, 7.0, 8.0]); @@ -231,4 +1378,785 @@ mod tests { let got = m.read(4).await; assert!(got.as_tensor().equal(qt.as_tensor())); } + + #[tokio::test] + async fn non_identity_skew_moves_physical_banks_and_roundtrips() { + let ty = f32_plain(); + let m = MatrixSram::with_banks(4, 16, 1, ty); + let view = MatrixLayout { + rows: 4, + cols: 4, + tile_count: 1, + tile_pitch_rows: 4, + alpha: 1, + tile_skew: 0, + }; + let values = (0..16).map(|v| v as f32 + 1.0).collect::>(); + m.write_layout_tile(0, view, 0, tile(ty, &values)).await; + let row_major = MatrixLayout { alpha: 0, ..view }; + assert_ne!( + m.physical_coord(0, view, 0, 1, 0), + m.physical_coord(0, row_major, 0, 1, 0) + ); + let got = m.read_layout_tile(0, view, 0).await; + assert_eq!(tensor_to_f32_vec(got.as_tensor()), values); + } + + #[tokio::test] + async fn wrong_skew_returns_wrong_values_not_only_extra_cycles() { + let ty = f32_plain(); + let m = MatrixSram::with_banks(4, 16, 1, ty); + let placed = MatrixLayout { + rows: 4, + cols: 4, + tile_count: 1, + tile_pitch_rows: 4, + alpha: 1, + tile_skew: 0, + }; + let wrong = MatrixLayout { alpha: 0, ..placed }; + let values = (0..16).map(|v| v as f32 + 1.0).collect::>(); + m.write_layout_tile(0, placed, 0, tile(ty, &values)).await; + let got = tensor_to_f32_vec(m.read_layout_tile(0, wrong, 0).await.as_tensor()); + assert_ne!(got, values); + } + + #[tokio::test] + async fn pending_dma_fills_physical_words() { + let ty = f32_plain(); + let m = MatrixSram::with_banks(4, 16, 1, ty); + let pending = m.mark_pending_tiles(0, 1).await; + let values = (0..16).map(|v| v as f32 + 1.0).collect::>(); + let (tx, rx) = oneshot::channel(); + assert!(tx.send(tile(ty, &values)).is_ok()); + m.fill_pending(pending, rx).await; + assert_eq!(tensor_to_f32_vec(m.read(0).await.as_tensor()), values); + } + + #[tokio::test] + async fn pending_dma_fills_multiple_tiles() { + let ty = f32_plain(); + let m = MatrixSram::with_banks(4, 16, 1, ty); + let pending = m.mark_pending_tiles(0, 2).await; + let values = (0..32).map(|value| value as f32 + 1.0).collect::>(); + let (tx, rx) = oneshot::channel(); + assert!(tx.send(tile(ty, &values)).is_ok()); + m.fill_pending(pending, rx).await; + + assert_eq!(tensor_to_f32_vec(m.read(0).await.as_tensor()), values[..16]); + assert_eq!( + tensor_to_f32_vec(m.read(16).await.as_tensor()), + values[16..] + ); + } + + #[tokio::test] + async fn per_view_skew_removes_cross_tile_bank_conflict_with_same_values() { + let ty = f32_plain(); + let m = MatrixSram::with_banks(4, 16, 1, ty); + let row_major = MatrixLayout { + rows: 1, + cols: 1, + tile_count: 4, + tile_pitch_rows: 1, + alpha: 0, + tile_skew: 0, + }; + let affine = MatrixLayout { + alpha: 1, + ..row_major + }; + let values = vec![11.0, 22.0, 33.0, 44.0]; + + let row_write = m.write_layout_packet(0, row_major, tile(ty, &values)).await; + let (row_values, row_read) = m.read_layout_packet(0, row_major).await; + assert_eq!(tensor_to_f32_vec(row_values.as_tensor()), values); + assert_eq!((row_write.service_cycles, row_read.service_cycles), (4, 4)); + + let affine_base = 4 * 4; + let affine_write = m + .write_layout_packet(affine_base, affine, tile(ty, &values)) + .await; + let (affine_values, affine_read) = m.read_layout_packet(affine_base, affine).await; + assert_eq!(tensor_to_f32_vec(affine_values.as_tensor()), values); + assert_eq!( + (affine_write.service_cycles, affine_read.service_cycles), + (1, 1) + ); + assert_eq!(affine_read.bank_stall_cycles, 0); + } + + #[tokio::test] + async fn bank_word_aligned_base_supplies_a_constant_field_phase() { + let ty = f32_plain(); + let sram = MatrixSram::with_banks(16, 32, 4, ty); + let view = MatrixLayout { + rows: 1, + cols: 4, + tile_count: 1, + tile_pitch_rows: 1, + alpha: 1, + tile_skew: 0, + }; + let unphased = sram.physical_coord(0, view, 0, 0, 0); + let phased = sram.physical_coord(3 * sram.bank_width(), view, 0, 0, 0); + assert_eq!(unphased.bank_row, phased.bank_row); + assert_eq!(phased.bank, (unphased.bank + 3) % sram.banks()); + + let left = tile(ty, &[1.0, 2.0, 3.0, 4.0]); + let right = tile(ty, &[5.0, 6.0, 7.0, 8.0]); + sram.write_layout_packet(0, view, left).await; + sram.write_layout_packet(3 * sram.bank_width(), view, right) + .await; + assert_eq!( + tensor_to_f32_vec(sram.read_layout_packet(0, view).await.0.as_tensor()), + [1.0, 2.0, 3.0, 4.0] + ); + assert_eq!( + tensor_to_f32_vec( + sram.read_layout_packet(3 * sram.bank_width(), view) + .await + .0 + .as_tensor() + ), + [5.0, 6.0, 7.0, 8.0] + ); + } + + #[tokio::test] + async fn diagonal_placement_serves_rows_and_columns_at_the_bank_floor() { + let ty = f32_plain(); + let values = (0..16).map(|value| value as f32 + 1.0).collect::>(); + let row_major = MatrixLayout { + rows: 4, + cols: 4, + tile_count: 1, + tile_pitch_rows: 4, + alpha: 0, + tile_skew: 0, + }; + let diagonal = MatrixLayout { + alpha: 1, + ..row_major + }; + + let row_sram = MatrixSram::with_banks(4, 16, 1, ty); + row_sram + .write_layout_tile(0, row_major, 0, tile(ty, &values)) + .await; + let (row, row_service) = row_sram + .read_layout_line(0, row_major, 0, 2, MatrixAccessAxis::Row) + .await; + let (column, column_service) = row_sram + .read_layout_line(0, row_major, 0, 1, MatrixAccessAxis::Column) + .await; + assert_eq!( + tensor_to_f32_vec(row.as_tensor()), + vec![9.0, 10.0, 11.0, 12.0] + ); + assert_eq!( + tensor_to_f32_vec(column.as_tensor()), + vec![2.0, 6.0, 10.0, 14.0] + ); + assert_eq!(row_service.service_cycles, 1); + assert_eq!(column_service.service_cycles, 4); + assert_eq!(column_service.bank_stall_cycles, 3); + + let diagonal_sram = MatrixSram::with_banks(4, 16, 1, ty); + diagonal_sram + .write_layout_tile(0, diagonal, 0, tile(ty, &values)) + .await; + let (row, row_service) = diagonal_sram + .read_layout_line(0, diagonal, 0, 2, MatrixAccessAxis::Row) + .await; + let (column, column_service) = diagonal_sram + .read_layout_line(0, diagonal, 0, 1, MatrixAccessAxis::Column) + .await; + assert_eq!( + tensor_to_f32_vec(row.as_tensor()), + vec![9.0, 10.0, 11.0, 12.0] + ); + assert_eq!( + tensor_to_f32_vec(column.as_tensor()), + vec![2.0, 6.0, 10.0, 14.0] + ); + assert_eq!(row_service.service_cycles, 1); + assert_eq!(column_service.service_cycles, 1); + assert_eq!(column_service.bank_stall_cycles, 0); + + let (by_row, _) = diagonal_sram + .read_layout_tile_axis(0, diagonal, 0, MatrixAccessAxis::Row) + .await; + let (by_column, _) = diagonal_sram + .read_layout_tile_axis(0, diagonal, 0, MatrixAccessAxis::Column) + .await; + assert_eq!(tensor_to_f32_vec(by_row.as_tensor()), values); + assert_eq!(tensor_to_f32_vec(by_column.as_tensor()), values); + } + + #[tokio::test] + async fn column_reads_restore_every_lane_for_real_bank_widths() { + const MLEN: u32 = 32; + const ROWS: u32 = 8; + let ty = bf16_plain(); + let values = (0..ROWS * MLEN) + .map(|index| index as f32) + .collect::>(); + + for bank_width in [1, 4, 32] { + let banks = MLEN / bank_width; + let sram = MatrixSram::with_banks(MLEN, 64, bank_width, ty); + let view = MatrixLayout { + rows: ROWS, + cols: MLEN, + tile_count: 1, + tile_pitch_rows: ROWS, + alpha: u32::from(banks > 1), + tile_skew: 0, + }; + sram.write_layout_tile(0, view, 0, tile(ty, &values)).await; + + for col in 0..MLEN { + let (line, _) = sram + .read_layout_line(0, view, 0, col, MatrixAccessAxis::Column) + .await; + let expected = (0..ROWS) + .map(|row| values[(row * MLEN + col) as usize]) + .collect::>(); + assert_eq!( + tensor_to_f32_vec(line.as_tensor()), + expected, + "column {col} failed at bank_width={bank_width}" + ); + } + } + } + + #[tokio::test] + async fn f32_matrix_words_roundtrip_with_multiple_lanes() { + const MLEN: u32 = 16; + let ty = f32_plain(); + let sram = MatrixSram::with_banks(MLEN, 32, 4, ty); + let view = MatrixLayout { + rows: 4, + cols: MLEN, + tile_count: 1, + tile_pitch_rows: 4, + alpha: 1, + tile_skew: 0, + }; + let values = (0..4 * MLEN).map(|index| index as f32).collect::>(); + sram.write_layout_tile(0, view, 0, tile(ty, &values)).await; + let got = sram.read_layout_tile(0, view, 0).await; + assert_eq!(tensor_to_f32_vec(got.as_tensor()), values); + } + + #[tokio::test] + async fn kda_prefill_state_becomes_decode_state_by_column_view_not_transpose_copy() { + const DIM: u32 = 8; + let ty = f32_plain(); + let sram = MatrixSram::with_banks(DIM, 64, 1, ty); + let view = MatrixLayout { + rows: DIM, + cols: DIM, + tile_count: 1, + tile_pitch_rows: DIM, + alpha: 1, + tile_skew: 0, + }; + // Logical storage is [value, key]. It is intentionally non-symmetric, + // because Kimi's 128x128 real shape cannot detect an axis error by shape. + let prefill = (0..DIM) + .flat_map(|value| (0..DIM).map(move |key| (value * 100 + key * 3 + 1) as f32)) + .collect::>(); + sram.write_layout_tile(0, view, 0, tile(ty, &prefill)).await; + + let mut decode = Vec::with_capacity((DIM * DIM) as usize); + let mut total_service = 0; + for key in 0..DIM { + let (line, service) = sram + .read_layout_line(0, view, 0, key, MatrixAccessAxis::Column) + .await; + assert_eq!(service.service_cycles, service.ideal_cycles); + total_service += service.service_cycles; + decode.extend(tensor_to_f32_vec(line.as_tensor())); + } + let expected = (0..DIM) + .flat_map(|key| (0..DIM).map(move |value| (value * 100 + key * 3 + 1) as f32)) + .collect::>(); + assert_eq!(decode, expected); + assert_ne!(prefill, expected, "row/column mismatch must be observable"); + assert_eq!(total_service, u64::from(DIM)); + } + + fn official_recurrent_layout( + rows: u32, + cols: u32, + tiles: u32, + bank_width: u32, + ) -> MatrixLayout { + let words_per_row = cols / bank_width; + MatrixLayout { + rows, + cols, + tile_count: tiles, + // Both the tile phase and the phase seen by an equal-row packet + // must traverse every bank-word group. Twice the row width is the + // smallest pitch that satisfies both constraints for the official + // power-of-two Mamba/KDA shapes. + tile_pitch_rows: 2 * words_per_row, + alpha: 1, + tile_skew: words_per_row, + } + } + + async fn official_recurrent_group_roundtrip(cols: u32, tiles: u32, expected_last_row: u32) { + const MLEN: u32 = 2048; + const DEPTH_ROWS: usize = 256; + const STATE_ROWS: u32 = 128; + const BF16_VALUES_PER_BANK_WORD: u32 = 32; + + let ty = bf16_plain(); + let sram = MatrixSram::with_banks(MLEN, DEPTH_ROWS, BF16_VALUES_PER_BANK_WORD, ty); + let view = official_recurrent_layout(STATE_ROWS, cols, tiles, BF16_VALUES_PER_BANK_WORD); + let final_row = (tiles - 1) * view.tile_pitch_rows + STATE_ROWS; + assert_eq!(final_row, expected_last_row); + assert!(final_row <= DEPTH_ROWS as u32); + + let value_count = tiles * STATE_ROWS * cols; + let values = (0..value_count) + .map(|index| ((index % 257) as f32 - 128.0) / 64.0) + .collect::>(); + let input = QuantTensor::quantize(Tensor::from_slice(&values), ty); + let expected = tensor_to_f32_vec(input.as_tensor()); + let write = sram.write_layout_packet(0, view, input).await; + assert_eq!(write.bank_stall_cycles, 0); + + // This is the recurrence access: the same logical state row from all + // heads in the group must fill all 64 bank words exactly once. + for row in [0, STATE_ROWS / 2, STATE_ROWS - 1] { + let lines = (0..tiles).map(|tile| (tile, row)).collect::>(); + let (packet, service) = sram.read_layout_indexed_rows(0, view, &lines).await; + assert_eq!(service.ideal_cycles, 1); + assert_eq!(service.service_cycles, 1); + assert_eq!(service.bank_stall_cycles, 0); + let got = tensor_to_f32_vec(packet.as_tensor()); + let expected_row = (0..tiles) + .flat_map(|tile| { + let start = ((tile * STATE_ROWS + row) * cols) as usize; + expected[start..start + cols as usize].iter().copied() + }) + .collect::>(); + assert_eq!(got, expected_row); + } + + let (roundtrip, read) = sram.read_layout_packet(0, view).await; + assert_eq!(tensor_to_f32_vec(roundtrip.as_tensor()), expected); + assert_eq!(read.bank_stall_cycles, 0); + + // A *single-base descriptor* with no per-tile phase can hold only two + // full-height heads. The stronger fixed-wiring D' control is tested + // separately below and must not be confused with this ISA limitation. + let single_descriptor_required_rows = STATE_ROWS * tiles; + assert!(single_descriptor_required_rows > DEPTH_ROWS as u32); + assert_eq!(DEPTH_ROWS as u32 / STATE_ROWS, 2); + } + + async fn fixed_phased_official_state_roundtrip(cols: u32, tiles: u32) { + const MLEN: u32 = 2048; + const DEPTH_ROWS: usize = 256; + const STATE_ROWS: u32 = 128; + const BF16_VALUES_PER_BANK_WORD: u32 = 32; + + let ty = bf16_plain(); + let sram = MatrixSram::with_banks(MLEN, DEPTH_ROWS, BF16_VALUES_PER_BANK_WORD, ty); + let words_per_head = cols / BF16_VALUES_PER_BANK_WORD; + assert_eq!(tiles * words_per_head, 64); + let fixed = MatrixLayout { + rows: STATE_ROWS, + cols, + tile_count: 1, + tile_pitch_rows: STATE_ROWS, + alpha: 1, + tile_skew: 0, + }; + let affine = MatrixLayout { + rows: STATE_ROWS, + cols, + tile_count: tiles, + tile_pitch_rows: 0, + alpha: 1, + tile_skew: words_per_head, + }; + + let names = (0..tiles) + .map(|tile| format!("head_{tile}")) + .collect::>(); + let bases = (0..tiles) + .map(|tile| tile * words_per_head * BF16_VALUES_PER_BANK_WORD) + .collect::>(); + let views = names + .iter() + .zip(&bases) + .map(|(name, &base)| (name.as_str(), base, fixed)) + .collect::>(); + sram.validate_disjoint_layouts(&views) + .expect("fixed per-head phases must fit without aliases"); + + // D' and D occupy exactly the same bank word for every official state + // value. D uses one compact descriptor; D' uses one ordinary base per + // head. Therefore programmable skew has no pure bank-service credit. + for tile in 0..tiles { + for row in 0..STATE_ROWS { + for word in 0..words_per_head { + let col = word * BF16_VALUES_PER_BANK_WORD; + let fixed_coord = + sram.physical_coord_unchecked(bases[tile as usize], fixed, 0, row, col); + let affine_coord = sram.physical_coord_unchecked(0, affine, tile, row, col); + assert_eq!(fixed_coord, affine_coord); + } + } + } + + let mut expected = Vec::with_capacity(tiles as usize); + for tile in 0..tiles { + let value_count = (STATE_ROWS * cols) as usize; + let values = (0..value_count) + .map(|index| { + let code = (tile as usize * 131 + index) % 257; + (code as f32 - 128.0) / 64.0 + }) + .collect::>(); + let input = QuantTensor::quantize(Tensor::from_slice(&values), ty); + expected.push(tensor_to_f32_vec(input.as_tensor())); + let write = sram + .write_layout_packet(bases[tile as usize], fixed, input) + .await; + assert_eq!(write.bank_stall_cycles, 0); + } + + sram.reset_packet_counters(); + let requests = bases.iter().map(|&base| (base, fixed)).collect::>(); + let (actual, service) = sram.read_layout_packets(&requests).await; + assert_eq!(service.ideal_cycles, u64::from(STATE_ROWS)); + assert_eq!(service.service_cycles, u64::from(STATE_ROWS)); + assert_eq!(service.bank_stall_cycles, 0); + assert_eq!(service.values, u64::from(tiles * STATE_ROWS * cols)); + for (packet, expected) in actual.iter().zip(expected) { + assert_eq!(tensor_to_f32_vec(packet.as_tensor()), expected); + } + } + + fn paper_addr(row: u32, bank_phase: u32) -> u32 { + const MLEN: u32 = 2048; + const BLEN: u32 = 32; + row * MLEN + bank_phase * BLEN + } + + fn expected_after_sram_quantization(sram: &MatrixSram, values: &[f32]) -> Vec { + values + .chunks(sram.bank_width as usize) + .flat_map(|word| { + let bytes = sram.values_to_word_bytes(word); + sram.word_bytes_to_values(&bytes) + .into_iter() + .take(word.len()) + }) + .collect() + } + + async fn assert_colocated_views_roundtrip( + sram: &MatrixSram, + views: &[(&str, u32, MatrixLayout)], + ) { + sram.validate_disjoint_layouts(views) + .expect("compiler placement must be physically disjoint"); + let mut expected = Vec::with_capacity(views.len()); + for (view_index, &(_, base, layout)) in views.iter().enumerate() { + let count = (layout.rows * layout.cols * layout.tile_count) as usize; + let values = (0..count) + .map(|index| (view_index * 100_000 + index) as f32) + .collect::>(); + let quantized = expected_after_sram_quantization(sram, &values); + let input = QuantTensor::quantize(Tensor::from_slice(&values), sram.ty()); + sram.write_layout_packet(base, layout, input).await; + expected.push(quantized); + } + for ((name, base, layout), expected) in views.iter().zip(expected) { + let (output, _) = sram.read_layout_packet(*base, *layout).await; + assert_eq!( + tensor_to_f32_vec(output.as_tensor()), + expected, + "co-resident view {name} did not round-trip", + ); + } + } + + fn kda_chunk_views(affine: bool) -> Vec<(&'static str, u32, MatrixLayout)> { + let (state, scalar, vector, scalar_row, vector_row) = if affine { + ( + MatrixLayout { + rows: 16, + cols: 128, + tile_count: 16, + tile_pitch_rows: 8, + alpha: 1, + tile_skew: 4, + }, + MatrixLayout { + rows: 16, + cols: 32, + tile_count: 16, + tile_pitch_rows: 1, + alpha: 1, + tile_skew: 3, + }, + MatrixLayout { + rows: 1, + cols: 128, + tile_count: 16, + tile_pitch_rows: 1, + alpha: 1, + tile_skew: 3, + }, + 136, + 168, + ) + } else { + ( + MatrixLayout { + rows: 16, + cols: 128, + tile_count: 16, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + MatrixLayout { + rows: 16, + cols: 32, + tile_count: 16, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + MatrixLayout { + rows: 1, + cols: 128, + tile_count: 16, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + 0, + 0, + ) + }; + vec![ + ("state", paper_addr(0, 0), state), + ( + "decay", + paper_addr(scalar_row, if affine { 0 } else { 4 }), + scalar, + ), + ( + "key", + paper_addr(scalar_row, if affine { 1 } else { 5 }), + scalar, + ), + ( + "query", + paper_addr(scalar_row, if affine { 2 } else { 6 }), + scalar, + ), + ( + "value_or_error", + paper_addr(vector_row, if affine { 0 } else { 8 }), + vector, + ), + ( + "prediction_or_output", + paper_addr(vector_row, if affine { 4 } else { 12 }), + vector, + ), + ] + } + + fn mamba_chunk_views(affine: bool) -> Vec<(&'static str, u32, MatrixLayout)> { + let tiles = if affine { 32 } else { 16 }; + let (state, scalar, vector, scalar_row, vector_row) = if affine { + ( + MatrixLayout { + rows: 16, + cols: 64, + tile_count: tiles, + tile_pitch_rows: 4, + alpha: 1, + tile_skew: 2, + }, + MatrixLayout { + rows: 16, + cols: 32, + tile_count: tiles, + tile_pitch_rows: 1, + alpha: 1, + tile_skew: 1, + }, + MatrixLayout { + rows: 1, + cols: 64, + tile_count: tiles, + tile_pitch_rows: 1, + alpha: 1, + tile_skew: 1, + }, + 140, + 188, + ) + } else { + ( + MatrixLayout { + rows: 16, + cols: 64, + tile_count: tiles, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + MatrixLayout { + rows: 16, + cols: 32, + tile_count: tiles, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + MatrixLayout { + rows: 1, + cols: 64, + tile_count: tiles, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }, + 0, + 0, + ) + }; + vec![ + ("state", paper_addr(0, 0), state), + ( + "decay_and_b", + paper_addr(scalar_row, if affine { 0 } else { 2 }), + scalar, + ), + ( + "c", + paper_addr(scalar_row, if affine { 16 } else { 3 }), + scalar, + ), + ( + "dt_and_skip", + paper_addr(scalar_row, if affine { 32 } else { 4 }), + scalar, + ), + ( + "x", + paper_addr(vector_row, if affine { 0 } else { 8 }), + vector, + ), + ( + "scratch", + paper_addr(vector_row, if affine { 2 } else { 10 }), + vector, + ), + ( + "output", + paper_addr(vector_row, if affine { 4 } else { 12 }), + vector, + ), + ] + } + + #[tokio::test] + async fn official_nemotron_state_group_fits_and_reads_32_heads_without_conflict() { + // 32 heads x [128,64] BF16 state. The compact affine layout occupies + // 252 physical rows; a fixed full-height layout can hold only 2 heads. + official_recurrent_group_roundtrip(64, 32, 252).await; + } + + #[tokio::test] + async fn official_kimi_state_group_fits_and_reads_16_heads_without_conflict() { + // 16 heads x [128,128] BF16 state. The compact affine layout occupies + // 248 physical rows; a fixed full-height layout can hold only 2 heads. + official_recurrent_group_roundtrip(128, 16, 248).await; + } + + #[tokio::test] + async fn fixed_phased_nemotron_state_matches_affine_without_programmable_skew() { + fixed_phased_official_state_roundtrip(64, 32).await; + } + + #[tokio::test] + async fn fixed_phased_kimi_state_matches_affine_without_programmable_skew() { + fixed_phased_official_state_roundtrip(128, 16).await; + } + + #[tokio::test] + async fn official_kda_chunk_and_all_live_fields_share_one_matrix_sram() { + const MLEN: u32 = 2048; + const DEPTH_ROWS: usize = 256; + const BLEN: u32 = 32; + for affine in [false, true] { + let sram = MatrixSram::with_banks(MLEN, DEPTH_ROWS, BLEN, bf16_plain()); + let views = kda_chunk_views(affine); + assert_colocated_views_roundtrip(&sram, &views).await; + let state = views[0]; + let lines = (0..state.2.tile_count) + .map(|tile| (tile, 0)) + .collect::>(); + let (_, service) = sram + .read_layout_indexed_rows(state.1, state.2, &lines) + .await; + assert_eq!(service.service_cycles, if affine { 1 } else { 4 }); + } + } + + #[tokio::test] + async fn official_mamba_chunk_and_all_live_fields_share_one_matrix_sram() { + const MLEN: u32 = 2048; + const DEPTH_ROWS: usize = 256; + const BLEN: u32 = 32; + for affine in [false, true] { + let sram = MatrixSram::with_banks(MLEN, DEPTH_ROWS, BLEN, bf16_plain()); + let views = mamba_chunk_views(affine); + assert_colocated_views_roundtrip(&sram, &views).await; + let state = views[0]; + let lines = (0..state.2.tile_count) + .map(|tile| (tile, 0)) + .collect::>(); + let (_, service) = sram + .read_layout_indexed_rows(state.1, state.2, &lines) + .await; + assert_eq!(service.service_cycles, if affine { 1 } else { 4 }); + } + } + + #[test] + fn colocated_view_validator_rejects_cross_tensor_aliases() { + let sram = MatrixSram::with_banks(2048, 256, 32, bf16_plain()); + let view = MatrixLayout { + rows: 1, + cols: 128, + tile_count: 16, + tile_pitch_rows: 16, + alpha: 1, + tile_skew: 0, + }; + let error = sram + .validate_disjoint_layouts(&[("first", 0, view), ("second", 0, view)]) + .unwrap_err(); + assert!(error.contains("first")); + assert!(error.contains("second")); + } } diff --git a/transactional_emulator/src/accelerator/access.rs b/transactional_emulator/src/accelerator/access.rs index 4f6239b3..bf162559 100644 --- a/transactional_emulator/src/accelerator/access.rs +++ b/transactional_emulator/src/accelerator/access.rs @@ -83,10 +83,11 @@ pub(crate) enum Cfg { Stride, VMask, TopkPolicy, + MatrixView, } impl Cfg { - pub(crate) const COUNT: usize = 4; + pub(crate) const COUNT: usize = 5; pub(crate) fn index(self) -> usize { match self { @@ -94,6 +95,7 @@ impl Cfg { Cfg::Stride => 1, Cfg::VMask => 2, Cfg::TopkPolicy => 3, + Cfg::MatrixView => 4, } } } @@ -237,57 +239,81 @@ pub(crate) fn op_access( op::Opcode::Invalid => OpAccess::none(Unit::Scalar), // === Matrix accumulate ops === - op::Opcode::M_MM { rs1, rs2 } | op::Opcode::M_TMM { rs1, rs2 } => OpAccess::new( - Unit::Matrix, - vec![ + op::Opcode::M_MM { rs1, rs2, view } | op::Opcode::M_TMM { rs1, rs2, view } => { + let mut reads = vec![ Gp(rs1), Gp(rs2), matrix_tile_at(gp(rs1)), vector(gp(rs2), *MLEN * *BLEN), Accum(AccumKind::M), - ], - vec![Accum(AccumKind::M)], - ), - op::Opcode::M_BMM { rs1, rs2 } | op::Opcode::M_BTMM { rs1, rs2 } => OpAccess::new( - Unit::Matrix, - vec![ + ]; + if view.is_some() { + reads.push(Resource::Cfg(Cfg::MatrixView)); + } + OpAccess::new(Unit::Matrix, reads, vec![Accum(AccumKind::M)]) + } + op::Opcode::M_BMM { rs1, rs2, view } | op::Opcode::M_BTMM { rs1, rs2, view } => { + let mut reads = vec![ Gp(rs1), Gp(rs2), matrix_tile_at(gp(rs1)), vector(gp(rs2), matrix_tile), Accum(AccumKind::Hm), - ], - vec![Accum(AccumKind::Hm)], - ), - op::Opcode::M_MV { rs1, rs2 } | op::Opcode::M_TMV { rs1, rs2 } => OpAccess::new( - Unit::Matrix, - vec![ + ]; + if view.is_some() { + reads.push(Resource::Cfg(Cfg::MatrixView)); + } + OpAccess::new(Unit::Matrix, reads, vec![Accum(AccumKind::Hm)]) + } + op::Opcode::M_MV { rs1, rs2, view } | op::Opcode::M_TMV { rs1, rs2, view } => { + let mut reads = vec![ Gp(rs1), Gp(rs2), matrix_tile_at(gp(rs1)), vector(gp(rs2), vector_tile), Accum(AccumKind::V), - ], - vec![Accum(AccumKind::V)], - ), - op::Opcode::M_BMV { rs1, rs2, rd } | op::Opcode::M_BTMV { rs1, rs2, rd } => OpAccess::new( - Unit::Matrix, - vec![ + ]; + if view.is_some() { + reads.push(Resource::Cfg(Cfg::MatrixView)); + } + OpAccess::new(Unit::Matrix, reads, vec![Accum(AccumKind::V)]) + } + op::Opcode::M_BMV { rs1, rs2, rd, view } | op::Opcode::M_BTMV { rs1, rs2, rd, view } => { + let mut reads = vec![ Gp(rs1), Gp(rs2), Gp(rd), matrix_tile_at(gp(rs1).wrapping_add(gp(rd))), vector(gp(rs2), vector_tile), Accum(AccumKind::Hv), - ], - vec![Accum(AccumKind::Hv)], - ), + ]; + if view.is_some() { + reads.push(Resource::Cfg(Cfg::MatrixView)); + } + OpAccess::new(Unit::Matrix, reads, vec![Accum(AccumKind::Hv)]) + } // === Matrix write-outs === // `mm_wo` is a read-modify-write: for each of `blen` rows it reads // `vec_base + i * mlen * stride_len`, splices the accumulator in, and // writes the row back. - op::Opcode::M_MM_WO { rd, rstride, imm } => { + op::Opcode::M_MM_WO { + rd, + rstride, + imm, + view, + } => { + if view.is_some() { + let mut reads = vec![Gp(rd), Accum(AccumKind::M), Resource::Cfg(Cfg::MatrixView)]; + if rstride != 0 { + reads.push(Gp(rstride)); + } + return OpAccess::new( + Unit::Matrix, + reads, + vec![Accum(AccumKind::M), matrix_tile_at(gp(rd))], + ); + } let stride_len = if rstride == 0 { 1 } else { gp(rstride) }; let base = row_base(gp(rd).wrapping_add(imm)); let span = (*BLEN) @@ -344,18 +370,21 @@ pub(crate) fn op_access( rs1, rs2, rmask, + .. } | op::Opcode::V_SUB_VV { rd, rs1, rs2, rmask, + .. } | op::Opcode::V_MUL_VV { rd, rs1, rs2, rmask, + .. } => { let mut reads = vec![ Gp(rd), @@ -550,6 +579,18 @@ pub(crate) fn op_access( ], vec![vector(gp(rd), *VLEN * *PREFETCH_V_AMOUNT)], ), + op::Opcode::H_PREFETCH_V_MV { rd, rs1, rs2, .. } => OpAccess::new( + Unit::Dma, + vec![ + Gp(rd), + Gp(rs1), + Hbm(rs2), + Resource::Cfg(Cfg::Scale), + Resource::Cfg(Cfg::Stride), + Resource::Cfg(Cfg::MatrixView), + ], + vec![matrix(gp(rd), matrix_tile)], + ), // A store reads the vram region it drains. Its HBM-side write is not // tracked as a resource; dispatch conservatively drains all pending // prefetches before an H_STORE_V instead (HBM WAR/RAW). @@ -566,6 +607,19 @@ pub(crate) fn op_access( vec![], ), + op::Opcode::H_STORE_V_MV { rd, rs1, rs2, .. } => OpAccess::new( + Unit::Dma, + vec![ + Gp(rd), + Gp(rs1), + Hbm(rs2), + Resource::Cfg(Cfg::Scale), + Resource::Cfg(Cfg::Stride), + Resource::Cfg(Cfg::MatrixView), + matrix(gp(rd), matrix_tile), + ], + vec![], + ), // === Control === op::Opcode::C_SET_ADDR_REG { rd, rs1, rs2 } => { OpAccess::new(Unit::Scalar, vec![Gp(rs1), Gp(rs2)], vec![Hbm(rd)]) @@ -584,6 +638,24 @@ pub(crate) fn op_access( vec![Gp(rd)], vec![Resource::Cfg(Cfg::TopkPolicy)], ), + op::Opcode::L_TILE_CFG { shape, mapping, .. } => OpAccess::new( + Unit::Scalar, + vec![Gp(shape), Gp(mapping)], + vec![Resource::Cfg(Cfg::MatrixView)], + ), + op::Opcode::L_TILE_EXEC { rd, rs1, rs2, .. } => OpAccess::new( + Unit::Vector, + vec![ + Gp(rd), + Gp(rs1), + Gp(rs2), + Resource::Cfg(Cfg::MatrixView), + matrix(gp(rd), matrix_tile), + matrix(gp(rs1), matrix_tile), + matrix(gp(rs2), matrix_tile), + ], + vec![matrix(gp(rd), matrix_tile)], + ), op::Opcode::C_LOOP_START { rd, .. } => OpAccess::new(Unit::Scalar, vec![], vec![Gp(rd)]), op::Opcode::C_LOOP_END { rd } => OpAccess::new(Unit::Scalar, vec![Gp(rd)], vec![Gp(rd)]), // C_BREAK writes the innermost loop's counter register, which is only @@ -623,7 +695,11 @@ mod tests { #[test] fn matrix_multiply_reads_regs_tile_batch_and_accumulator() { - let a = access(op::Opcode::M_MM { rs1: 1, rs2: 2 }); + let a = access(op::Opcode::M_MM { + rs1: 1, + rs2: 2, + view: None, + }); assert_eq!(a.unit, Unit::Matrix); assert!(a.reads.contains(&Resource::Gp(1))); assert!(a.reads.contains(&Resource::Gp(2))); @@ -646,6 +722,7 @@ mod tests { rd: 2, rstride: 0, imm: 0, + view: None, }); let expected_span = (*BLEN - 1) * *MLEN + *VLEN; assert_eq!( @@ -668,6 +745,7 @@ mod tests { rd: 2, rstride: 3, imm: 0, + view: None, }); assert!(a.reads.contains(&Resource::Gp(3))); let expected_span = (*BLEN - 1) * *MLEN * gp_stub(3) + *VLEN; @@ -708,6 +786,7 @@ mod tests { rs1: 2, rs2: 3, rmask: 0, + view_mask: 0, }); assert_eq!(a.unit, Unit::Vector); assert_eq!( @@ -729,6 +808,7 @@ mod tests { rs1: 2, rs2: 3, rmask: 1, + view_mask: 0, }); assert!(masked.reads.contains(&Resource::Cfg(Cfg::VMask))); } diff --git a/transactional_emulator/src/accelerator/dispatch.rs b/transactional_emulator/src/accelerator/dispatch.rs index 02af2e8f..855b8a33 100644 --- a/transactional_emulator/src/accelerator/dispatch.rs +++ b/transactional_emulator/src/accelerator/dispatch.rs @@ -4,21 +4,24 @@ //! match and dispatch-only helpers. use half::bf16; -use quantize::MxDataType; +use quantize::{MxDataType, QuantTensor}; use crate::runtime_config::PERIOD; use crate::runtime_config::{ HLEN, MATRIX_KV_TYPE, MATRIX_WEIGHT_TYPE, MLEN, PREFETCH_M_AMOUNT, PREFETCH_V_AMOUNT, SCALAR_FP_BASIC_CYCLES, SCALAR_FP_EXP_CYCLES, SCALAR_FP_RECI_CYCLES, SCALAR_FP_SQRT_CYCLES, - SCALAR_INT_BASIC_CYCLES, STORE_V_AMOUNT, VECTOR_ACTIVATION_TYPE, VECTOR_KV_TYPE, VLEN, + SCALAR_INT_BASIC_CYCLES, STATE_TYPE, STORE_V_AMOUNT, VECTOR_ACTIVATION_TYPE, VECTOR_KV_TYPE, + VLEN, }; use crate::stage_profile::{ResourceKind, StageProfiler}; +use crate::vector_machine::VectorBinaryOp; use crate::{cycle, dma, op, timing}; use runtime::{Executor, Instant}; use super::Accelerator; use super::access::{self, OpAccess}; use super::loop_state::LoopDecision; +use super::matrix_view::{LTileExecArgs, MatrixViewBinaryArgs}; use super::scoreboard::{DmaKind, Scoreboard}; /// How `do_ops` charges time. @@ -95,11 +98,17 @@ impl Accelerator { // - H_PREFETCH_*: all outstanding stores (HBM RAW) plus // prefetches overlapping its SRAM destination (WAW). // - Everything else: SRAM overlaps only. - let pending = if access.barrier || matches!(op, op::Opcode::H_STORE_V { .. }) { + let pending = if access.barrier + || matches!( + op, + op::Opcode::H_STORE_V { .. } | op::Opcode::H_STORE_V_MV { .. } + ) { scoreboard.take_all_dma() } else if matches!( op, - op::Opcode::H_PREFETCH_M { .. } | op::Opcode::H_PREFETCH_V { .. } + op::Opcode::H_PREFETCH_M { .. } + | op::Opcode::H_PREFETCH_V { .. } + | op::Opcode::H_PREFETCH_V_MV { .. } ) { scoreboard.take_dma_for_prefetch(&access) } else { @@ -178,41 +187,73 @@ impl Accelerator { panic!("invalid opcode at pc {pc}"); } - op::Opcode::M_MM { rs1, rs2 } => { + op::Opcode::M_MM { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine - .mm(self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2)) + .mm_with_view( + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + view, + ) .await; } - op::Opcode::M_MM_WO { rd, rstride, imm } => { + op::Opcode::M_MM_WO { + rd, + rstride, + imm, + view, + } => { let stride_len = if *rstride == 0 { 1 } else { self.reg_file.read_gp(*rstride) }; - self.m_machine - .mm_wo(self.reg_file.read_gp(*rd) + *imm, stride_len) - .await; + if let Some(view) = self.resolve_matrix_view(*view, pc) { + let logical_offset = if *rstride == 0 { + *imm + } else { + self.reg_file.read_gp(*rstride).wrapping_add(*imm) + }; + let service = self + .m_machine + .mview_wo(self.reg_file.read_gp(*rd), logical_offset, view) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + } else { + self.m_machine + .mm_wo(self.reg_file.read_gp(*rd) + *imm, stride_len) + .await; + } } - op::Opcode::M_TMM { rs1, rs2 } => { + op::Opcode::M_TMM { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine - .tmm(self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2)) + .tmm( + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + view, + ) .await; } - op::Opcode::M_BMM { rs1, rs2 } => { + op::Opcode::M_BMM { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine .bmm( self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2), self.reg_file.bmm_scale(), + view, ) .await; } - op::Opcode::M_BTMM { rs1, rs2 } => { + op::Opcode::M_BTMM { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine .btmm( self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2), self.reg_file.bmm_scale(), + view, ) .await; } @@ -221,31 +262,45 @@ impl Accelerator { .bmm_wo(self.reg_file.read_gp(*rd) + *imm) .await; } - op::Opcode::M_MV { rs1, rs2 } => { + op::Opcode::M_MV { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine - .mv(self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2)) + .mv( + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + view, + ) .await; } - op::Opcode::M_TMV { rs1, rs2 } => { + op::Opcode::M_TMV { rs1, rs2, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine - .tmv(self.reg_file.read_gp(*rs1), self.reg_file.read_gp(*rs2)) + .tmv( + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + view, + ) .await; } - op::Opcode::M_BMV { rs1, rs2, rd } => { + op::Opcode::M_BMV { rs1, rs2, rd, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine .bmv( self.reg_file.read_gp(*rs1) + self.reg_file.read_gp(*rd), self.reg_file.read_gp(*rs2), self.reg_file.bmm_scale(), + view, ) .await; } - op::Opcode::M_BTMV { rs1, rs2, rd } => { + op::Opcode::M_BTMV { rs1, rs2, rd, view } => { + let view = self.resolve_matrix_view(*view, pc); self.m_machine .btmv( self.reg_file.read_gp(*rs1) + self.reg_file.read_gp(*rd), self.reg_file.read_gp(*rs2), self.reg_file.bmm_scale(), + view, ) .await; } @@ -265,17 +320,32 @@ impl Accelerator { rs1, rs2, rmask, + view_mask, } => { let mask = self.resolve_v_mask(*rmask); - self.v_machine - .add( - self.reg_file.read_gp(*rd), - self.reg_file.read_gp(*rs1), - self.reg_file.read_gp(*rs2), - *rmask, + if Self::matrix_view_operand_mask(*view_mask).is_some() { + self.vector_binary_with_matrix_views(MatrixViewBinaryArgs { + operation: VectorBinaryOp::Add, + rd: *rd, + rs1: *rs1, + rs2: *rs2, + rmask: *rmask, + view_mask: *view_mask, mask, - ) + pc, + }) .await; + } else { + self.v_machine + .add( + self.reg_file.read_gp(*rd), + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + *rmask, + mask, + ) + .await; + } } op::Opcode::V_ADD_VF { rd, @@ -299,17 +369,32 @@ impl Accelerator { rs1, rs2, rmask, + view_mask, } => { let mask = self.resolve_v_mask(*rmask); - self.v_machine - .sub( - self.reg_file.read_gp(*rd), - self.reg_file.read_gp(*rs1), - self.reg_file.read_gp(*rs2), - *rmask, + if Self::matrix_view_operand_mask(*view_mask).is_some() { + self.vector_binary_with_matrix_views(MatrixViewBinaryArgs { + operation: VectorBinaryOp::Sub, + rd: *rd, + rs1: *rs1, + rs2: *rs2, + rmask: *rmask, + view_mask: *view_mask, mask, - ) + pc, + }) .await; + } else { + self.v_machine + .sub( + self.reg_file.read_gp(*rd), + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + *rmask, + mask, + ) + .await; + } } op::Opcode::V_SUB_VF { rd, @@ -335,17 +420,32 @@ impl Accelerator { rs1, rs2, rmask, + view_mask, } => { let mask = self.resolve_v_mask(*rmask); - self.v_machine - .mul( - self.reg_file.read_gp(*rd), - self.reg_file.read_gp(*rs1), - self.reg_file.read_gp(*rs2), - *rmask, + if Self::matrix_view_operand_mask(*view_mask).is_some() { + self.vector_binary_with_matrix_views(MatrixViewBinaryArgs { + operation: VectorBinaryOp::Mul, + rd: *rd, + rs1: *rs1, + rs2: *rs2, + rmask: *rmask, + view_mask: *view_mask, mask, - ) + pc, + }) .await; + } else { + self.v_machine + .mul( + self.reg_file.read_gp(*rd), + self.reg_file.read_gp(*rs1), + self.reg_file.read_gp(*rs2), + *rmask, + mask, + ) + .await; + } } op::Opcode::V_MUL_VF { rd, @@ -650,6 +750,7 @@ impl Accelerator { let dtype = match precision { op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, }; let region = self.mx_region(dtype, addr, offset, *rstride); @@ -667,6 +768,58 @@ impl Accelerator { .vram .continous_write_delayed(dest, *PREFETCH_V_AMOUNT, xfer) .await; + // SRAM bank write service follows completed DMA; it cannot overlap the fill. + } + op::Opcode::H_PREFETCH_V_MV { + rd, + rs1, + rs2, + rstride, + precision, + view, + } => { + let descriptor = self.resolve_matrix_view(Some(*view), pc).unwrap(); + let values = descriptor.values(); + let dtype = match precision { + op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, + op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, + }; + let region = self.mx_region( + dtype, + self.reg_file.read_hbm(*rs2), + self.reg_file.read_gp(*rs1), + *rstride, + ); + let xfer = dma::transfer_mx_from_hbm( + &self.hbm, + region, + self.m_machine.mram.ty(), + *VLEN, + values.div_ceil(*VLEN), + 1, + ); + let tensor = xfer.await.unwrap_or_else(|error| { + panic!("Matrix-view DMA receiver dropped: {error}") + }); + let tensor = if tensor.as_tensor().numel() == values as usize { + tensor + } else { + QuantTensor::quantize( + tensor.as_tensor().narrow(0, 0, i64::from(values)), + self.m_machine.mram.ty(), + ) + }; + let service = self + .m_machine + .mram + .write_layout_packet( + self.reg_file.read_gp(*rd), + descriptor.layout(), + tensor, + ) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; } op::Opcode::H_STORE_V { rd, @@ -681,6 +834,7 @@ impl Accelerator { let dtype = match precision { op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, }; let region = self.mx_region(dtype, addr, offset, *rstride); @@ -695,6 +849,35 @@ impl Accelerator { ) .await; } + op::Opcode::H_STORE_V_MV { + rd, + rs1, + rs2, + rstride, + precision, + view, + } => { + let descriptor = self.resolve_matrix_view(Some(*view), pc).unwrap(); + let (packet, service) = self + .m_machine + .mram + .read_layout_packet(self.reg_file.read_gp(*rd), descriptor.layout()) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + let rows = dma::split_packet_rows(&packet, *VLEN, self.m_machine.mram.ty()); + let dtype = match precision { + op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, + op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, + }; + let region = self.mx_region( + dtype, + self.reg_file.read_hbm(*rs2), + self.reg_file.read_gp(*rs1), + *rstride, + ); + dma::store_rows_to_hbm(&self.hbm, region, rows, *VLEN).await; + } op::Opcode::C_SET_ADDR_REG { rd, rs1, rs2 } => { let imm = ((self.reg_file.read_gp(*rs1) as u64) << 32) | (self.reg_file.read_gp(*rs2) as u64); @@ -717,6 +900,40 @@ impl Accelerator { self.reg_file.set_topk_policy(self.reg_file.read_gp(*rd)); cycle!(1); } + op::Opcode::L_TILE_CFG { + shape, + mapping, + slot, + } => { + self.reg_file + .configure_mview(*slot, *shape, *mapping) + .unwrap_or_else(|error| { + tracing::error!(pc, slot, %error, "invalid L_TILE_CFG"); + panic!("{error} at pc {pc}") + }); + cycle!(1); + } + op::Opcode::L_TILE_EXEC { + rd, + rs1, + rs2, + primitive, + source_axis, + scale_axis, + } => { + self.execute_l_tile( + LTileExecArgs { + destination_register: *rd, + source_register: *rs1, + scale_register: *rs2, + primitive: *primitive, + source_axis: *source_axis, + scale_axis: *scale_axis, + }, + pc, + ) + .await; + } op::Opcode::C_LOOP_START { rd, imm } => { self.loop_state.start(pc, *rd, *imm, &mut self.reg_file); cycle!(1); @@ -754,6 +971,8 @@ impl Accelerator { op::Opcode::H_PREFETCH_M { .. } | op::Opcode::H_PREFETCH_V { .. } | op::Opcode::H_STORE_V { .. } + | op::Opcode::H_PREFETCH_V_MV { .. } + | op::Opcode::H_STORE_V_MV { .. } ); if after > issue && !is_dma_op { tracing::warn!( @@ -827,6 +1046,102 @@ impl Accelerator { } fn op_access_for_opcode(&self, op: &op::Opcode) -> OpAccess { + if let op::Opcode::H_PREFETCH_V_MV { view, .. } | op::Opcode::H_STORE_V_MV { view, .. } = op + { + let descriptor = self.reg_file.matrix_view(*view).unwrap_or_else(|error| { + panic!("{error} while building Matrix-view DMA scoreboard access") + }); + let mut dma_access = access::op_access(op, &|reg| self.reg_file.read_gp(reg), &|| { + self.reg_file.topk_policy() + }); + for resource in dma_access + .reads + .iter_mut() + .chain(dma_access.writes.iter_mut()) + { + if let access::Resource::Sram(range) = resource + && range.space == access::SramSpace::Matrix + { + range.len = descriptor.values(); + } + } + return dma_access; + } + + let matrix_vector = match op { + op::Opcode::V_ADD_VV { + rd, + rs1, + rs2, + rmask, + view_mask, + } + | op::Opcode::V_SUB_VV { + rd, + rs1, + rs2, + rmask, + view_mask, + } + | op::Opcode::V_MUL_VV { + rd, + rs1, + rs2, + rmask, + view_mask, + } if Self::matrix_view_operand_mask(*view_mask).is_some() => { + Some((*rd, *rs1, *rs2, *rmask, *view_mask)) + } + _ => None, + }; + if let Some((rd, rs1, rs2, rmask, encoded_mask)) = matrix_vector { + use access::{Cfg, Resource, SramRange, SramSpace, Unit}; + + let view_mask = Self::matrix_view_operand_mask(encoded_mask).unwrap(); + let mut reads = vec![ + Resource::Gp(rd), + Resource::Gp(rs1), + Resource::Gp(rs2), + Resource::Cfg(Cfg::MatrixView), + ]; + if rmask != 0 { + reads.push(Resource::Cfg(Cfg::VMask)); + } + for (slot, register) in [(1_u8, rs1), (2_u8, rs2)] { + let (space, len) = if view_mask & (1 << slot) != 0 { + let view = self.reg_file.matrix_view(slot).unwrap_or_else(|error| { + panic!("{error} while building Matrix-view scoreboard access") + }); + (SramSpace::Matrix, view.values()) + } else { + (SramSpace::Vector, *VLEN) + }; + reads.push(Resource::Sram(SramRange::new( + space, + self.reg_file.read_gp(register), + len, + ))); + } + let (space, len) = if view_mask & 0b001 != 0 { + let view = self.reg_file.matrix_view(0).unwrap_or_else(|error| { + panic!("{error} while building Matrix-view scoreboard access") + }); + (SramSpace::Matrix, view.values()) + } else { + (SramSpace::Vector, *VLEN) + }; + return OpAccess { + unit: Unit::Vector, + barrier: false, + reads, + writes: vec![Resource::Sram(SramRange::new( + space, + self.reg_file.read_gp(rd), + len, + ))], + }; + } + access::op_access(op, &|reg| self.reg_file.read_gp(reg), &|| { self.reg_file.topk_policy() }) @@ -896,6 +1211,7 @@ impl Accelerator { let dtype = match precision { op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, }; let region = self.mx_region(dtype, addr, offset, *rstride); let xfer = dma::transfer_mx_from_hbm( @@ -920,6 +1236,60 @@ impl Accelerator { }); (DmaKind::Prefetch, done_rx) } + op::Opcode::H_PREFETCH_V_MV { + rd, + rs1, + rs2, + rstride, + precision, + view, + } => { + let descriptor = self + .reg_file + .matrix_view(*view) + .unwrap_or_else(|error| panic!("{error} while issuing Matrix-view prefetch")); + let values = descriptor.values(); + let dtype = match precision { + op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, + op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, + }; + let region = self.mx_region( + dtype, + self.reg_file.read_hbm(*rs2), + self.reg_file.read_gp(*rs1), + *rstride, + ); + let xfer = dma::transfer_mx_from_hbm( + &self.hbm, + region, + self.m_machine.mram.ty(), + *VLEN, + values.div_ceil(*VLEN), + 1, + ); + let dest = self.reg_file.read_gp(*rd); + let (pending, service) = self + .m_machine + .mram + .mark_pending_layout_packet(dest, descriptor.layout()) + .await; + let mram = self.m_machine.mram.clone(); + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + Executor::current().spawn(async move { + let tensor = xfer.await.unwrap_or_else(|error| { + panic!("Matrix-view DMA receiver dropped: {error}") + }); + Executor::current() + .resolve_at(PERIOD * service.service_cycles.max(1)) + .await; + let (tx, rx) = tokio::sync::oneshot::channel(); + let _ = tx.send(tensor); + mram.fill_pending(pending, rx).await; + let _ = done_tx.send(Executor::current().now()); + }); + (DmaKind::Prefetch, done_rx) + } op::Opcode::H_STORE_V { rd, rs1, @@ -933,6 +1303,7 @@ impl Accelerator { let dtype = match precision { op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, }; let region = self.mx_region(dtype, addr, offset, *rstride); // Snapshot the source rows at issue: later instructions may @@ -948,6 +1319,46 @@ impl Accelerator { }); (DmaKind::Store, done_rx) } + op::Opcode::H_STORE_V_MV { + rd, + rs1, + rs2, + rstride, + precision, + view, + } => { + let descriptor = self + .reg_file + .matrix_view(*view) + .unwrap_or_else(|error| panic!("{error} while issuing Matrix-view store")); + let (packet, service) = self + .m_machine + .mram + .read_layout_packet(self.reg_file.read_gp(*rd), descriptor.layout()) + .await; + let rows = dma::split_packet_rows(&packet, *VLEN, self.m_machine.mram.ty()); + let dtype = match precision { + op::VectorPrecision::Activation => *VECTOR_ACTIVATION_TYPE, + op::VectorPrecision::KeyValue => *VECTOR_KV_TYPE, + op::VectorPrecision::State => *STATE_TYPE, + }; + let region = self.mx_region( + dtype, + self.reg_file.read_hbm(*rs2), + self.reg_file.read_gp(*rs1), + *rstride, + ); + let hbm = self.hbm.clone(); + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + Executor::current().spawn(async move { + Executor::current() + .resolve_at(PERIOD * service.service_cycles.max(1)) + .await; + dma::store_rows_to_hbm(&hbm, region, rows, *VLEN).await; + let _ = done_tx.send(Executor::current().now()); + }); + (DmaKind::Store, done_rx) + } _ => return false, }; let writes: Vec = access @@ -978,7 +1389,8 @@ fn resource_kind_for_opcode(op: &op::Opcode) -> ResourceKind { | op::Opcode::M_MV_WO { .. } | op::Opcode::M_BMV_WO { .. } => ResourceKind::Matrix, - op::Opcode::V_ADD_VV { .. } + op::Opcode::L_TILE_EXEC { .. } + | op::Opcode::V_ADD_VV { .. } | op::Opcode::V_ADD_VF { .. } | op::Opcode::V_SUB_VV { .. } | op::Opcode::V_SUB_VF { .. } @@ -1015,13 +1427,16 @@ fn resource_kind_for_opcode(op: &op::Opcode) -> ResourceKind { | op::Opcode::C_SET_STRIDE_REG { .. } | op::Opcode::C_SET_V_MASK_REG { .. } | op::Opcode::C_SET_TOPK_REG { .. } + | op::Opcode::L_TILE_CFG { .. } | op::Opcode::C_LOOP_START { .. } | op::Opcode::C_LOOP_END { .. } | op::Opcode::C_BREAK => ResourceKind::Scalar, op::Opcode::H_PREFETCH_M { .. } | op::Opcode::H_PREFETCH_V { .. } - | op::Opcode::H_STORE_V { .. } => ResourceKind::Dma, + | op::Opcode::H_STORE_V { .. } + | op::Opcode::H_PREFETCH_V_MV { .. } + | op::Opcode::H_STORE_V_MV { .. } => ResourceKind::Dma, op::Opcode::Invalid => ResourceKind::Other, } diff --git a/transactional_emulator/src/accelerator/matrix_view.rs b/transactional_emulator/src/accelerator/matrix_view.rs new file mode 100644 index 00000000..a95755ed --- /dev/null +++ b/transactional_emulator/src/accelerator/matrix_view.rs @@ -0,0 +1,659 @@ +//! Compiler-programmable views over PLENA's fixed-diagonal Matrix SRAM. +//! +//! A view is architectural placement metadata, not a cache or a traversal +//! engine. Existing Matrix operations name one of four slots explicitly. +//! There is no implicit selection, replacement, auto-advance, or model state. + +use super::Accelerator; +use crate::vector_machine::{TileScaleLayout, VectorBinaryOp}; +use crate::{op, timing}; +use quantize::{QuantTensor, tensor_from_f32_slice, tensor_to_f32_vec}; +use sram::matrix::MatrixLayout; +use sram::matrix::MatrixPacketService; + +const VIEW_SLOTS: usize = 4; +const DIM_MASK: u32 = (1 << 12) - 1; +const TILE_COUNT_MASK: u32 = (1 << 8) - 1; +const PITCH_MASK: u32 = (1 << 16) - 1; +const PHASE_MASK: u32 = (1 << 6) - 1; +const BROADCAST_MINOR: u8 = 1 << 3; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct MatrixViewShape { + pub(crate) rows: u32, + pub(crate) cols: u32, + pub(crate) tile_count: u32, +} + +impl MatrixViewShape { + pub(crate) fn unpack(word: u32) -> Self { + Self { + rows: (word & DIM_MASK) + 1, + cols: ((word >> 12) & DIM_MASK) + 1, + tile_count: ((word >> 24) & TILE_COUNT_MASK) + 1, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct MatrixViewMap { + /// Distance between consecutive logical tiles, measured in physical rows. + pub(crate) tile_pitch_rows: u32, + /// Compiler-selected bank phase stride between consecutive logical tiles. + pub(crate) tile_phase_stride: u32, + pub(crate) flags: u8, +} + +impl MatrixViewMap { + pub(crate) fn unpack(word: u32) -> Result { + if (word >> 16) & PHASE_MASK != 0 { + return Err("Matrix-view mapping bits [21:16] are reserved".into()); + } + let mapping = Self { + tile_pitch_rows: word & PITCH_MASK, + tile_phase_stride: (word >> 22) & PHASE_MASK, + flags: ((word >> 28) & 0xf) as u8, + }; + if mapping.flags & !BROADCAST_MINOR != 0 { + return Err(format!( + "Matrix-view flags contain reserved bits: {:#x}", + mapping.flags + )); + } + Ok(mapping) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct MatrixViewDescriptor { + pub(crate) shape: MatrixViewShape, + pub(crate) mapping: MatrixViewMap, +} + +impl MatrixViewDescriptor { + fn unpack(shape_word: u32, map_word: u32) -> Result { + Ok(Self { + shape: MatrixViewShape::unpack(shape_word), + mapping: MatrixViewMap::unpack(map_word)?, + }) + } + + fn validate(self, banks: u32, bank_width: u32) -> Result { + if !banks.is_power_of_two() || banks > 64 { + return Err(format!( + "Matrix-view bank count must be a power of two in 1..=64, got {banks}" + )); + } + if bank_width == 0 { + return Err("Matrix-view bank width must be positive".into()); + } + if !self.shape.cols.is_multiple_of(bank_width) { + return Err(format!( + "Matrix-view width {} is not a multiple of bank width {bank_width}", + self.shape.cols + )); + } + let words_per_row = self.shape.cols / bank_width; + let row_groups = words_per_row.div_ceil(banks); + let alpha = 1; + let tile_phase_stride = self.mapping.tile_phase_stride; + let mut occupied = std::collections::HashMap::new(); + for tile in 0..self.shape.tile_count { + for row in 0..self.shape.rows { + for word in 0..words_per_row { + let bank_row = + tile * self.mapping.tile_pitch_rows + row * row_groups + word / banks; + let bank = (alpha * bank_row + tile_phase_stride * tile + word) % banks; + if let Some(previous) = occupied.insert((bank, bank_row), (tile, row, word)) { + return Err(format!( + "Matrix view aliases logical bank words: {previous:?} and {:?} at bank={bank}, row={bank_row}", + (tile, row, word) + )); + } + } + } + } + Ok(self) + } + + pub(crate) fn layout(self) -> MatrixLayout { + MatrixLayout { + rows: self.shape.rows, + cols: self.shape.cols, + tile_count: self.shape.tile_count, + tile_pitch_rows: self.mapping.tile_pitch_rows, + // The row term always uses PLENA's prior-work diagonal wiring. + // Only the inter-tile phase is compiler selected. + alpha: 1, + tile_skew: self.mapping.tile_phase_stride, + } + } + + pub(crate) fn broadcast_minor(self) -> bool { + self.mapping.flags & BROADCAST_MINOR != 0 + } + + pub(crate) fn values(self) -> u32 { + self.shape + .rows + .checked_mul(self.shape.cols) + .and_then(|value| value.checked_mul(self.shape.tile_count)) + .expect("validated Matrix-view dimensions overflowed u32") + } +} + +pub(super) struct MatrixViewTable { + banks: u32, + bank_width: u32, + slots: [Option; VIEW_SLOTS], +} + +impl MatrixViewTable { + pub(super) fn new(banks: u32, bank_width: u32) -> Self { + assert!(banks.is_power_of_two()); + assert!(banks <= 64); + assert!(bank_width > 0); + Self { + banks, + bank_width, + slots: [None; VIEW_SLOTS], + } + } + + pub(super) fn configure( + &mut self, + slot: u8, + shape_word: u32, + map_word: u32, + ) -> Result<(), String> { + let index = self.slot_index(slot)?; + let descriptor = MatrixViewDescriptor::unpack(shape_word, map_word)? + .validate(self.banks, self.bank_width)?; + self.slots[index] = Some(descriptor); + Ok(()) + } + + pub(super) fn get(&self, slot: u8) -> Result { + let index = self.slot_index(slot)?; + self.slots[index].ok_or_else(|| format!("Matrix-view slot {slot} is not configured")) + } + + fn slot_index(&self, slot: u8) -> Result { + let index = usize::from(slot); + if index >= VIEW_SLOTS { + Err(format!( + "Matrix-view slot {slot} is outside 0..{VIEW_SLOTS}" + )) + } else { + Ok(index) + } + } +} + +pub(super) struct MatrixViewBinaryArgs { + pub(super) operation: VectorBinaryOp, + pub(super) rd: u8, + pub(super) rs1: u8, + pub(super) rs2: u8, + pub(super) rmask: u8, + pub(super) view_mask: u8, + pub(super) mask: u32, + pub(super) pc: usize, +} + +pub(super) struct LTileExecArgs { + pub(super) destination_register: u8, + pub(super) source_register: u8, + pub(super) scale_register: u8, + pub(super) primitive: op::LTilePrimitive, + pub(super) source_axis: op::LTileAxis, + pub(super) scale_axis: op::LTileAxis, +} + +impl Accelerator { + pub(super) fn resolve_matrix_view( + &self, + slot: Option, + pc: usize, + ) -> Option { + slot.map(|slot| { + self.reg_file.matrix_view(slot).unwrap_or_else(|error| { + tracing::error!(pc, slot, %error, "invalid Matrix-view consumer"); + panic!("{error} at pc {pc}") + }) + }) + } + + pub(super) async fn read_l_tile_lines( + &mut self, + base: u32, + view: MatrixViewDescriptor, + axis: op::LTileAxis, + lines: &[(u32, u32)], + ) -> (QuantTensor, MatrixPacketService) { + match axis { + op::LTileAxis::Row => { + self.m_machine + .mram + .read_layout_indexed_rows(base, view.layout(), lines) + .await + } + op::LTileAxis::Column => { + self.m_machine + .mram + .read_layout_indexed_columns(base, view.layout(), lines) + .await + } + } + } + + /// Decode the explicit Matrix-view operand marker carried by the VV + /// family. Bits 0/1/2 select destination/source-1/source-2 slots. Keeping + /// the marker in the instruction avoids inferring addressing semantics + /// from whichever configuration registers happen to be live. + pub(super) fn matrix_view_operand_mask(view_mask: u8) -> Option { + assert!(view_mask < 8, "Matrix-view operand mask exceeds three bits"); + (view_mask != 0).then_some(view_mask) + } + + pub(super) async fn vector_binary_with_matrix_views(&mut self, args: MatrixViewBinaryArgs) { + let MatrixViewBinaryArgs { + operation, + rd, + rs1, + rs2, + rmask, + view_mask, + mask, + pc, + } = args; + let view_mask = Self::matrix_view_operand_mask(view_mask) + .expect("Matrix-view vector helper requires the explicit marker"); + assert_ne!(view_mask, 0, "Matrix-view operand mask cannot be zero"); + + let destination_view = + (view_mask & 0b001 != 0).then(|| self.resolve_matrix_view(Some(0), pc).unwrap()); + let source1_view = + (view_mask & 0b010 != 0).then(|| self.resolve_matrix_view(Some(1), pc).unwrap()); + let source2_view = + (view_mask & 0b100 != 0).then(|| self.resolve_matrix_view(Some(2), pc).unwrap()); + + for descriptor in [destination_view, source1_view, source2_view] + .into_iter() + .flatten() + { + assert_eq!( + descriptor.values(), + self.v_machine.tile_size(), + "a Vector Matrix-view operand must restore exactly VLEN values" + ); + } + + let mut requests = Vec::with_capacity(2); + if let Some(view) = source1_view { + requests.push((self.reg_file.read_gp(rs1), view.layout())); + } + if let Some(view) = source2_view { + requests.push((self.reg_file.read_gp(rs2), view.layout())); + } + let matrix_packets = if requests.is_empty() { + Vec::new() + } else { + let (packets, service) = self.m_machine.mram.read_layout_packets(&requests).await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + packets + }; + let mut matrix_packets = matrix_packets.into_iter(); + let lhs = if source1_view.is_some() { + matrix_packets + .next() + .expect("missing Matrix source-1 packet") + } else { + self.v_machine.vram.read(self.reg_file.read_gp(rs1)).await + }; + let rhs = if source2_view.is_some() { + matrix_packets + .next() + .expect("missing Matrix source-2 packet") + } else { + self.v_machine.vram.read(self.reg_file.read_gp(rs2)).await + }; + debug_assert!(matrix_packets.next().is_none()); + + let result = self + .v_machine + .binary_packet(operation, lhs, rhs, rmask, mask) + .await; + if let Some(view) = destination_view { + let service = self + .m_machine + .mram + .write_layout_packet(self.reg_file.read_gp(rd), view.layout(), result) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + } else { + self.v_machine + .vram + .write(self.reg_file.read_gp(rd), result) + .await; + } + } + + /// Execute one model-independent recurrence primitive over Matrix views. + /// + /// Views 0/1/2 are destination/source/scalars. The decoder owns only a + /// deterministic row/column walk; all bases, shapes and layouts remain + /// compiler-visible architectural state. Matrix-view storage is BF16. + pub(super) async fn execute_l_tile(&mut self, args: LTileExecArgs, pc: usize) { + let LTileExecArgs { + destination_register, + source_register, + scale_register, + primitive, + source_axis, + scale_axis, + } = args; + let destination = self.resolve_matrix_view(Some(0), pc).unwrap(); + let source = self.resolve_matrix_view(Some(1), pc).unwrap(); + let scales = self.resolve_matrix_view(Some(2), pc).unwrap(); + let dst_base = self.reg_file.read_gp(destination_register); + let src_base = self.reg_file.read_gp(source_register); + let scale_base = self.reg_file.read_gp(scale_register); + + if !scales.broadcast_minor() { + panic!("L_TILE scale view must set BROADCAST_MINOR"); + } + if scales.shape.tile_count != 1 && scales.shape.tile_count != destination.shape.tile_count { + panic!("L_TILE scale tiles must be one or match destination tiles"); + } + if source.shape.tile_count != 1 && source.shape.tile_count != destination.shape.tile_count { + panic!("L_TILE source tiles must be one or match destination tiles"); + } + + let dst_layout = destination.layout(); + let source_line_count = match source_axis { + op::LTileAxis::Row => source.shape.rows, + op::LTileAxis::Column => source.shape.cols, + }; + let source_line_width = match source_axis { + op::LTileAxis::Row => source.shape.cols, + op::LTileAxis::Column => source.shape.rows, + }; + let scale_line_count = match scale_axis { + op::LTileAxis::Row => scales.shape.rows, + op::LTileAxis::Column => scales.shape.cols, + }; + let scale_line_width = match scale_axis { + op::LTileAxis::Row => scales.shape.cols, + op::LTileAxis::Column => scales.shape.rows, + }; + + // A recurrence line is serviced by one existing Vector operation. + // Wider Matrix views remain legal for DMA/Matrix consumers, but this + // controller does not split a logical line across multiple VLEN ops. + assert!( + source_line_width <= self.v_machine.tile_size(), + "L_TILE logical line width exceeds VLEN; compiler must tile the columns" + ); + + match primitive { + op::LTilePrimitive::ScaleAccum | op::LTilePrimitive::OuterUpdate => { + if source_line_width != destination.shape.cols { + panic!("row-wise L_TILE source/destination widths differ"); + } + if source_line_count != 1 && source_line_count != destination.shape.rows { + panic!("row-wise L_TILE source rows must be one or match destination"); + } + if scale_line_count < destination.shape.rows { + panic!("L_TILE scale view has fewer logical lines than destination"); + } + let tiles_per_packet = (self.v_machine.tile_size() / destination.shape.cols).max(1); + for row in 0..destination.shape.rows { + for first_tile in + (0..destination.shape.tile_count).step_by(tiles_per_packet as usize) + { + let tile_count = + tiles_per_packet.min(destination.shape.tile_count - first_tile); + let scale_layout = if scales.shape.tile_count == 1 { + TileScaleLayout::Compact { first_tile } + } else { + TileScaleLayout::Expanded + }; + let destination_lines = (first_tile..first_tile + tile_count) + .map(|tile| (tile, row)) + .collect::>(); + let source_lines = destination_lines + .iter() + .map(|&(tile, destination_row)| { + ( + if source.shape.tile_count == 1 { + 0 + } else { + tile + }, + if source_line_count == 1 { + 0 + } else { + destination_row + }, + ) + }) + .collect::>(); + let scale_lines = if scales.shape.tile_count == 1 { + // Compact per-segment scalars are fetched once, + // one cycle ahead of the all-bank state packet. + vec![(0, row)] + } else { + destination_lines + .iter() + .map(|&(tile, destination_row)| (tile, destination_row)) + .collect::>() + }; + + let (dst_packet, dst_service) = self + .m_machine + .mram + .read_layout_indexed_rows(dst_base, dst_layout, &destination_lines) + .await; + timing::charge_bank_cycles(dst_service.service_cycles.max(1)).await; + let (src_packet, src_service) = self + .read_l_tile_lines(src_base, source, source_axis, &source_lines) + .await; + timing::charge_bank_cycles(src_service.service_cycles.max(1)).await; + let (scale_packet, scale_service) = self + .read_l_tile_lines(scale_base, scales, scale_axis, &scale_lines) + .await; + // Scalar bank words are deliberately charged separately: + // a full state packet already consumes every bank word. + timing::charge_bank_cycles(scale_service.service_cycles.max(1)).await; + + let result = match primitive { + op::LTilePrimitive::ScaleAccum => { + self.v_machine + .tile_scale_accum( + dst_packet, + src_packet, + scale_packet, + destination.shape.cols, + scale_line_width, + scale_layout, + ) + .await + } + op::LTilePrimitive::OuterUpdate => { + self.v_machine + .tile_outer_update( + dst_packet, + src_packet, + scale_packet, + destination.shape.cols, + scale_line_width, + scale_layout, + ) + .await + } + op::LTilePrimitive::DotReduce => unreachable!(), + }; + let service = self + .m_machine + .mram + .write_layout_indexed_rows( + dst_base, + dst_layout, + &destination_lines, + result, + ) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + } + } + } + op::LTilePrimitive::DotReduce => { + if destination.shape.rows != 1 + || destination.shape.cols != source_line_width + || destination.shape.tile_count != source.shape.tile_count + { + panic!("DOT_REDUCE destination must be one row per source tile"); + } + if scale_line_count < source_line_count { + panic!("DOT_REDUCE scale view has fewer lines than reduction rows"); + } + let tiles_per_packet = (self.v_machine.tile_size() / source_line_width).max(1); + for first_tile in (0..source.shape.tile_count).step_by(tiles_per_packet as usize) { + let tile_count = tiles_per_packet.min(source.shape.tile_count - first_tile); + let scale_layout = if scales.shape.tile_count == 1 { + TileScaleLayout::Compact { first_tile } + } else { + TileScaleLayout::Expanded + }; + let destination_lines = (first_tile..first_tile + tile_count) + .map(|tile| (tile, 0)) + .collect::>(); + let (destination_packet, destination_service) = self + .m_machine + .mram + .read_layout_indexed_rows(dst_base, dst_layout, &destination_lines) + .await; + timing::charge_bank_cycles(destination_service.service_cycles.max(1)).await; + let mut accumulator = tensor_to_f32_vec(destination_packet.as_tensor()); + assert_eq!(accumulator.len(), (tile_count * source_line_width) as usize); + for row in 0..source_line_count { + let source_lines = (first_tile..first_tile + tile_count) + .map(|tile| (tile, row)) + .collect::>(); + let scale_lines = if scales.shape.tile_count == 1 { + vec![(0, row)] + } else { + source_lines + .iter() + .map(|&(tile, source_row)| (tile, source_row)) + .collect::>() + }; + let (source_packet, source_service) = self + .read_l_tile_lines(src_base, source, source_axis, &source_lines) + .await; + timing::charge_bank_cycles(source_service.service_cycles.max(1)).await; + let (scale_packet, scale_service) = self + .read_l_tile_lines(scale_base, scales, scale_axis, &scale_lines) + .await; + timing::charge_bank_cycles(scale_service.service_cycles.max(1)).await; + self.v_machine + .tile_dot_accumulate( + &mut accumulator, + source_packet, + scale_packet, + source_line_width, + scale_line_width, + scale_layout, + ) + .await; + } + let result = QuantTensor::quantize( + tensor_from_f32_slice(&accumulator), + self.m_machine.mram.ty(), + ); + let service = self + .m_machine + .mram + .write_layout_indexed_rows(dst_base, dst_layout, &destination_lines, result) + .await; + timing::charge_bank_cycles(service.service_cycles.max(1)).await; + } + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn shape(rows: u32, cols: u32, tiles: u32) -> u32 { + (rows - 1) | ((cols - 1) << 12) | ((tiles - 1) << 24) + } + + fn mapping(pitch: u32) -> u32 { + pitch + } + + #[test] + fn v3_mapping_words_match_the_python_contract() { + assert_eq!(mapping(64), 0x0000_0040); + assert_eq!(mapping(0) | (4 << 22), 0x0100_0000); + assert_eq!(mapping(0) | (4 << 22) | (8 << 28), 0x8100_0000); + } + + #[test] + fn configuration_matches_the_python_contract() { + let mut table = MatrixViewTable::new(16, 4); + table.configure(2, shape(64, 64, 3), mapping(64)).unwrap(); + let view = table.get(2).unwrap(); + assert_eq!( + view.shape, + MatrixViewShape { + rows: 64, + cols: 64, + tile_count: 3 + } + ); + assert_eq!(view.mapping.tile_pitch_rows, 64); + assert_eq!(view.mapping.tile_phase_stride, 0); + } + + #[test] + fn rejects_aliasing_pitch_and_reserved_mapping_bits() { + let mut table = MatrixViewTable::new(16, 4); + assert!(table.configure(0, shape(64, 64, 2), mapping(63)).is_err()); + assert!( + table + .configure(0, shape(64, 64, 2), mapping(64) | (1 << 16)) + .is_err() + ); + let phased = mapping(64) | (5 << 22); + table.configure(0, shape(64, 64, 2), phased).unwrap(); + let view = table.get(0).unwrap(); + assert_eq!((view.layout().alpha, view.layout().tile_skew), (1, 5)); + + let programmable_row = mapping(64) | (3 << 16) | (5 << 22); + assert!( + table + .configure(0, shape(64, 64, 2), programmable_row) + .is_err() + ); + for reserved_flag in [1_u32, 2, 4] { + let invalid = mapping(64) | (5 << 22) | (reserved_flag << 28); + assert!(table.configure(0, shape(64, 64, 2), invalid).is_err()); + } + } + + #[test] + fn zero_pitch_is_legal_only_when_tile_phase_prevents_aliasing() { + let mut table = MatrixViewTable::new(64, 32); + assert!(table.configure(0, shape(128, 128, 8), mapping(0)).is_err()); + let compact = mapping(0) | (4 << 22); + table.configure(0, shape(128, 128, 8), compact).unwrap(); + let view = table.get(0).unwrap(); + assert_eq!(view.mapping.tile_pitch_rows, 0); + assert_eq!((view.layout().alpha, view.layout().tile_skew), (1, 4)); + } +} diff --git a/transactional_emulator/src/accelerator/matrix_view_tests.rs b/transactional_emulator/src/accelerator/matrix_view_tests.rs new file mode 100644 index 00000000..efaed0bc --- /dev/null +++ b/transactional_emulator/src/accelerator/matrix_view_tests.rs @@ -0,0 +1,2148 @@ +//! Matrix SRAM view, projection and L_TILE recurrence integration tests. +//! +//! Tests execute real banked storage and inspect state/output values, packet +//! service, and Serial/Scoreboard behavior using prepared BF16 operands. + +use std::sync::{Arc, Mutex}; + +use half::bf16; +use memory::{ErasedMemoryModel, MemoryBacked, NaiveTiming, WithTiming}; +use quantize::{DataType, FpType, MxDataType, QuantTensor, tensor_to_f32_vec}; +use runtime::{Executor, Instant}; +use sram::matrix::{MatrixLayout, MatrixPacketCounterSnapshot}; +use sram::{MatrixSram, VectorSram}; +use tch::Tensor; + +use super::{Accelerator, Scoreboard, TimingDriver}; +use crate::matrix_machine::MatrixMachine; +use crate::op; +use crate::runtime_config::{ + BLEN, BROADCAST_AMOUNT, HLEN, MATRIX_SRAM_TYPE, MLEN, PERIOD, VECTOR_SRAM_TYPE, VLEN, +}; +use crate::timing::{TimingMode, set_timing_mode}; +use crate::vector_machine::VectorMachine; + +const TOKENS: usize = 4; +const HBM_PACKET_STRIDE_BYTES: u32 = 8192; +const HBM_TEST_CAPACITY: usize = 1 << 20; + +fn set_gp(rd: u8, value: u32) -> op::Opcode { + op::Opcode::S_ADDI_INT { + rd, + rs1: 0, + imm: value, + } +} + +fn ltile_shape_word(rows: u32, cols: u32, tiles: u32) -> u32 { + (rows - 1) | ((cols - 1) << 12) | ((tiles - 1) << 24) +} + +fn ltile_map_word(pitch: u32, tile_phase_stride: Option, broadcast_minor: bool) -> u32 { + let mut flags = 0_u32; + let mut word = pitch; + if let Some(phase) = tile_phase_stride { + // Match the physical layout used to seed the SRAM: treatment changes + // only the per-tile phase, while retaining PLENA's diagonal row term. + word |= phase << 22; + } + if broadcast_minor { + flags |= 1 << 3; + } + word | (flags << 28) +} + +fn configure_ltile_view(ops: &mut Vec, slot: u8, shape: u32, mapping: u32) { + ops.extend([ + set_gp(10, shape), + set_gp(11, mapping), + op::Opcode::L_TILE_CFG { + shape: 10, + mapping: 11, + slot, + }, + ]); +} + +async fn execute_oversized_ltile_line(primitive: op::LTilePrimitive, axis: op::LTileAxis) { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + executor.spawn(async move { + let ty = MxDataType::Plain(DataType::Fp(FpType::BF16)); + let mram = Arc::new(MatrixSram::with_banks(64, 256, 4, ty)); + let vram = Arc::new(VectorSram::from_mx_type(64, 8, ty)); + let m_machine = MatrixMachine::new(mram, vram.clone(), 64, 16, 4, 4); + let v_machine = VectorMachine::new(vram, 64, 16); + let hbm: Arc = Arc::new(WithTiming::new( + NaiveTiming::preset_ddr4_2400p(4), + MemoryBacked::with_capacity(4096), + )); + let mut accelerator = Accelerator::new(m_machine, v_machine, hbm); + let (source_rows, source_cols) = match axis { + op::LTileAxis::Row => (1, 128), + op::LTileAxis::Column => (128, 4), + }; + let mut ops = Vec::new(); + configure_ltile_view(&mut ops, 0, ltile_shape_word(1, 128, 1), 0); + configure_ltile_view( + &mut ops, + 1, + ltile_shape_word(source_rows, source_cols, 1), + 0, + ); + configure_ltile_view( + &mut ops, + 2, + ltile_shape_word(4, 4, 1), + ltile_map_word(0, None, true), + ); + ops.push(op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive, + source_axis: axis, + scale_axis: op::LTileAxis::Row, + }); + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + }); + executor.enter(Instant::ETERNITY).await; +} + +#[tokio::test] +#[should_panic(expected = "L_TILE logical line width exceeds VLEN")] +async fn ltile_scale_accum_rejects_a_row_wider_than_vlen() { + execute_oversized_ltile_line(op::LTilePrimitive::ScaleAccum, op::LTileAxis::Row).await; +} + +#[tokio::test] +#[should_panic(expected = "L_TILE logical line width exceeds VLEN")] +async fn ltile_outer_update_rejects_a_row_wider_than_vlen() { + execute_oversized_ltile_line(op::LTilePrimitive::OuterUpdate, op::LTileAxis::Row).await; +} + +#[tokio::test] +#[should_panic(expected = "L_TILE logical line width exceeds VLEN")] +async fn ltile_dot_reduce_rejects_a_column_wider_than_vlen() { + execute_oversized_ltile_line(op::LTilePrimitive::DotReduce, op::LTileAxis::Column).await; +} + +async fn new_accelerator(mram: Arc, vram: Arc) -> Accelerator { + let m_machine = MatrixMachine::new(mram, vram.clone(), *MLEN, *HLEN, *BLEN, *BROADCAST_AMOUNT); + let v_machine = VectorMachine::new(vram, *VLEN, *HLEN); + let hbm: Arc = Arc::new(WithTiming::new( + NaiveTiming::preset_ddr4_2400p(4), + MemoryBacked::with_capacity(4096), + )); + Accelerator::new(m_machine, v_machine, hbm) +} + +fn new_accelerator_with_hbm( + mram: Arc, + vram: Arc, + image: &[u8], +) -> (Accelerator, Arc>) { + assert!(image.len() <= HBM_TEST_CAPACITY); + let m_machine = MatrixMachine::new(mram, vram.clone(), *MLEN, *HLEN, *BLEN, *BROADCAST_AMOUNT); + let v_machine = VectorMachine::new(vram, *VLEN, *HLEN); + let hbm = Arc::new(WithTiming::new( + NaiveTiming::preset_ddr4_2400p(4), + MemoryBacked::with_capacity(HBM_TEST_CAPACITY), + )); + hbm.data() + .with_data(|bytes| bytes[..image.len()].copy_from_slice(image)); + let erased: Arc = hbm.clone(); + (Accelerator::new(m_machine, v_machine, erased), hbm) +} + +fn hbm_packet_offset(region: u32) -> u32 { + region * HBM_PACKET_STRIDE_BYTES +} + +fn write_hbm_state_packet(image: &mut [u8], region: u32, values: &[f32]) { + let offset = hbm_packet_offset(region) as usize; + let mut packet = QuantTensor::quantize(Tensor::from_slice(values), full_state_type()); + let (bytes, scale_bytes) = packet.into_bytes(); + assert!(scale_bytes.is_empty()); + assert!(bytes.len() <= HBM_PACKET_STRIDE_BYTES as usize); + image[offset..offset + bytes.len()].copy_from_slice(&bytes); +} + +fn read_hbm_state_packet( + hbm: &WithTiming, + region: u32, + values: usize, +) -> Vec { + let offset = hbm_packet_offset(region) as usize; + let bytes_per_value = full_state_type().element_type().size_in_bits() as usize / 8; + let byte_len = values * bytes_per_value; + let mut bytes = vec![0_u8; byte_len]; + hbm.data().with_data(|image| { + bytes.copy_from_slice(&image[offset..offset + byte_len]); + }); + let packet = QuantTensor::from_bytes(&bytes, &[], values, full_state_type()); + tensor_to_f32_vec(packet.as_tensor()) +} + +#[allow(clippy::too_many_arguments)] +fn append_matrix_view_dma( + ops: &mut Vec, + load: bool, + matrix_base: u32, + hbm_region: u32, + rows: u32, + cols: u32, + affine: bool, + broadcast_minor: bool, +) { + const DMA_VIEW: u8 = 3; + configure_ltile_view( + ops, + DMA_VIEW, + ltile_shape_word(rows, cols, if broadcast_minor { 1 } else { FULL_TILES }), + full_map(rows, cols, affine, broadcast_minor), + ); + ops.extend([ + set_gp(12, matrix_base), + set_gp(13, hbm_packet_offset(hbm_region)), + ]); + if load { + ops.push(op::Opcode::H_PREFETCH_V_MV { + rd: 12, + rs1: 13, + rs2: 0, + rstride: 0, + precision: op::VectorPrecision::State, + view: DMA_VIEW, + }); + } else { + ops.push(op::Opcode::H_STORE_V_MV { + rd: 12, + rs1: 13, + rs2: 0, + rstride: 0, + precision: op::VectorPrecision::State, + view: DMA_VIEW, + }); + } +} + +fn assert_close(actual: &[f32], expected: &[f32]) { + assert_eq!(actual.len(), expected.len()); + for (index, (actual, expected)) in actual.iter().zip(expected).enumerate() { + let tolerance = 1e-2 + 1e-2 * expected.abs(); + assert!( + (actual - expected).abs() <= tolerance, + "value {index}: expected {expected}, got {actual}, tolerance {tolerance}" + ); + } +} + +#[derive(Debug)] +struct LTileResult { + state: Vec, + output: Vec, + matrix: MatrixPacketCounterSnapshot, +} + +async fn check_ltile_coefficient_packets(primitive: op::LTilePrimitive) { + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let results = Arc::new(Mutex::new(Vec::new())); + let task_results = results.clone(); + executor.spawn(async move { + // Three full single-tile packets; a full packet plus a one-tile tail; + // and two full multi-tile packets. Different coefficients per tile + // make the global compact offset observable in every case. + for (cols, tiles) in [(64_u32, 3_u32), (32, 3), (32, 4)] { + for compact in [true, false] { + for shared_source in [false, true] { + let is_dot = matches!(primitive, op::LTilePrimitive::DotReduce); + if is_dot && shared_source { + // DOT_REDUCE requires one source tile per output tile. + continue; + } + let is_scale = matches!(primitive, op::LTilePrimitive::ScaleAccum); + let dst_rows = if is_dot { 1 } else { 2 }; + let src_rows = if is_dot { 2 } else { 1 }; + let src_tiles = if shared_source { 1 } else { tiles }; + let scale_tiles = if compact { 1 } else { tiles }; + let per_tile = if is_scale { 2 } else { 1 }; + let scale_cols = if compact { + (per_tile * tiles).div_ceil(*BLEN) * *BLEN + } else { + *BLEN + }; + let make_layout = |rows, cols, tile_count| MatrixLayout { + rows, + cols, + tile_count, + tile_pitch_rows: 2, + alpha: 1, + tile_skew: 0, + }; + let dst_layout = make_layout(dst_rows, cols, tiles); + let src_layout = make_layout(src_rows, cols, src_tiles); + let scale_layout = make_layout(2, scale_cols, scale_tiles); + let source_value = |tile: u32, row: u32, col: u32| { + (1 + if shared_source { 0 } else { tile } + row + col % 3) as f32 + }; + let mut source = Vec::new(); + for tile in 0..src_tiles { + for row in 0..src_rows { + for col in 0..cols { + source.push(source_value(tile, row, col)); + } + } + } + let mut scales = vec![-16.0; (scale_tiles * 2 * scale_cols) as usize]; + for tile in 0..tiles { + for row in 0..2 { + let index = if compact { + row * scale_cols + per_tile * tile + } else { + (tile * 2 + row) * scale_cols + } as usize; + if is_scale { + scales[index] = 0.5; + scales[index + 1] = (tile + row + 1) as f32; + } else { + scales[index] = (tile + row + 1) as f32; + } + } + } + let destination = vec![1.0; (tiles * dst_rows * cols) as usize]; + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 64, + *BLEN, + full_state_type(), + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + let mut ops = Vec::new(); + for (slot, base, layout, values) in [ + (0, 0, dst_layout, &destination), + (1, 1024, src_layout, &source), + (2, 2048, scale_layout, &scales), + ] { + mram.write_layout_packet( + base, + layout, + QuantTensor::quantize(Tensor::from_slice(values), full_state_type()), + ) + .await; + configure_ltile_view( + &mut ops, + slot, + ltile_shape_word(layout.rows, layout.cols, layout.tile_count), + ltile_map_word(layout.tile_pitch_rows, None, slot == 2), + ); + } + ops.extend([ + set_gp(1, 0), + set_gp(2, 1024), + set_gp(3, 2048), + op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive, + source_axis: op::LTileAxis::Row, + scale_axis: op::LTileAxis::Row, + }, + ]); + let mut accelerator = new_accelerator(mram.clone(), vram).await; + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + let actual = tensor_to_f32_vec( + mram.read_layout_packet(0, dst_layout).await.0.as_tensor(), + ); + let mut expected = Vec::with_capacity(destination.len()); + for tile in 0..tiles { + for row in 0..dst_rows { + for col in 0..cols { + let value = match primitive { + op::LTilePrimitive::ScaleAccum => { + 0.5 + (tile + row + 1) as f32 * source_value(tile, 0, col) + } + op::LTilePrimitive::OuterUpdate => { + 1.0 + (tile + row + 1) as f32 * source_value(tile, 0, col) + } + op::LTilePrimitive::DotReduce => { + 1.0 + (0..src_rows) + .map(|r| { + (tile + r + 1) as f32 * source_value(tile, r, col) + }) + .sum::() + } + }; + expected.push(value); + } + } + } + task_results.lock().unwrap().push(( + format!( + "cols={cols}, tiles={tiles}, compact={compact}, shared={shared_source}" + ), + actual, + expected, + )); + } + } + } + }); + executor.enter(Instant::ETERNITY).await; + let results = results.lock().unwrap(); + let expected_cases = if matches!(primitive, op::LTilePrimitive::DotReduce) { + 6 + } else { + 12 + }; + assert_eq!(results.len(), expected_cases); + for (case, actual, expected) in results.iter() { + // All inputs and outputs are BF16-exact; no tolerance can hide a + // coefficient, tile, row, or broadcast selection error. + assert_eq!(actual, expected, "{case}"); + } +} + +#[tokio::test] +async fn l_tile_scale_accum_preserves_coefficient_layout_across_packets() { + check_ltile_coefficient_packets(op::LTilePrimitive::ScaleAccum).await; +} + +#[tokio::test] +async fn l_tile_dot_reduce_preserves_coefficient_layout_across_packets() { + check_ltile_coefficient_packets(op::LTilePrimitive::DotReduce).await; +} + +#[tokio::test] +async fn l_tile_outer_update_preserves_coefficient_layout_across_packets() { + check_ltile_coefficient_packets(op::LTilePrimitive::OuterUpdate).await; +} + +async fn run_ltile_primitives(tile_skew: Option) -> LTileResult { + const ROWS: u32 = 2; + const TILES: u32 = 8; + const COLS: u32 = 8; + const SCALE_COLS: u32 = 4; + const PITCH: u32 = 8; + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + + let state_layout = MatrixLayout { + rows: ROWS, + cols: COLS, + tile_count: TILES, + tile_pitch_rows: PITCH, + alpha: 1, + tile_skew: tile_skew.unwrap_or(0), + }; + let source_layout = MatrixLayout { + rows: 1, + cols: COLS, + tile_count: TILES, + tile_pitch_rows: PITCH, + alpha: 1, + tile_skew: tile_skew.unwrap_or(0), + }; + let scale_layout = MatrixLayout { + rows: ROWS, + cols: SCALE_COLS, + tile_count: TILES, + tile_pitch_rows: PITCH, + alpha: 1, + tile_skew: tile_skew.unwrap_or(0), + }; + let output_layout = MatrixLayout { + rows: 1, + cols: COLS, + tile_count: TILES, + tile_pitch_rows: PITCH, + alpha: 1, + tile_skew: tile_skew.unwrap_or(0), + }; + let state_base = 0; + let source_base = 64 * *MLEN; + let scale_base = 128 * *MLEN; + let output_base = 192 * *MLEN; + let state_input = (0..TILES * ROWS * COLS) + .map(|index| 0.25 + index as f32 / 16.0) + .collect::>(); + let source = (0..TILES * COLS) + .map(|index| 1.0 + index as f32 / 32.0) + .collect::>(); + let scales = (0..TILES) + .flat_map(|tile| { + (0..ROWS) + .flat_map(move |row| [0.5 + row as f32 / 32.0 + tile as f32 / 64.0, 0.25, 0.0, 0.0]) + }) + .collect::>(); + + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 64, + *BLEN, + *MATRIX_SRAM_TYPE, + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + mram.write_layout_packet( + state_base, + state_layout, + QuantTensor::quantize(Tensor::from_slice(&state_input), *MATRIX_SRAM_TYPE), + ) + .await; + mram.write_layout_packet( + source_base, + source_layout, + QuantTensor::quantize(Tensor::from_slice(&source), *MATRIX_SRAM_TYPE), + ) + .await; + mram.write_layout_packet( + scale_base, + scale_layout, + QuantTensor::quantize(Tensor::from_slice(&scales), *MATRIX_SRAM_TYPE), + ) + .await; + mram.write_layout_packet( + output_base, + output_layout, + QuantTensor::quantize( + Tensor::zeros( + [(TILES * COLS) as i64], + (tch::Kind::Float, tch::Device::Cpu), + ), + *MATRIX_SRAM_TYPE, + ), + ) + .await; + + let mut accelerator = new_accelerator(mram.clone(), vram).await; + let state_map = ltile_map_word(PITCH, tile_skew, false); + let scale_map = ltile_map_word(PITCH, tile_skew, true); + let source_map = ltile_map_word(PITCH, tile_skew, false); + let output_map = ltile_map_word(PITCH, tile_skew, false); + let mut ops = Vec::new(); + + configure_ltile_view(&mut ops, 0, ltile_shape_word(ROWS, COLS, TILES), state_map); + configure_ltile_view(&mut ops, 1, ltile_shape_word(1, COLS, TILES), source_map); + configure_ltile_view( + &mut ops, + 2, + ltile_shape_word(ROWS, SCALE_COLS, TILES), + scale_map, + ); + ops.extend([ + set_gp(1, state_base), + set_gp(2, source_base), + set_gp(3, scale_base), + op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive: op::LTilePrimitive::ScaleAccum, + source_axis: op::LTileAxis::Row, + scale_axis: op::LTileAxis::Row, + }, + ]); + + configure_ltile_view(&mut ops, 0, ltile_shape_word(1, COLS, TILES), output_map); + configure_ltile_view(&mut ops, 1, ltile_shape_word(ROWS, COLS, TILES), state_map); + configure_ltile_view( + &mut ops, + 2, + ltile_shape_word(ROWS, SCALE_COLS, TILES), + scale_map, + ); + ops.extend([ + set_gp(1, output_base), + set_gp(2, state_base), + set_gp(3, scale_base), + op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive: op::LTilePrimitive::DotReduce, + source_axis: op::LTileAxis::Row, + scale_axis: op::LTileAxis::Row, + }, + ]); + + configure_ltile_view(&mut ops, 0, ltile_shape_word(ROWS, COLS, TILES), state_map); + configure_ltile_view(&mut ops, 1, ltile_shape_word(1, COLS, TILES), source_map); + configure_ltile_view( + &mut ops, + 2, + ltile_shape_word(ROWS, SCALE_COLS, TILES), + scale_map, + ); + ops.extend([ + set_gp(1, state_base), + set_gp(2, source_base), + set_gp(3, scale_base), + op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive: op::LTilePrimitive::OuterUpdate, + source_axis: op::LTileAxis::Row, + scale_axis: op::LTileAxis::Row, + }, + ]); + + mram.reset_packet_counters(); + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + let matrix = accelerator.matrix_view_packet_counters(); + let state = tensor_to_f32_vec( + mram.read_layout_packet(state_base, state_layout) + .await + .0 + .as_tensor(), + ); + let output = tensor_to_f32_vec( + mram.read_layout_packet(output_base, output_layout) + .await + .0 + .as_tensor(), + ); + LTileResult { + state, + output, + matrix, + } +} + +#[tokio::test] +async fn l_tile_exec_runs_row_and_column_primitives_through_the_same_banks() { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + *task_result.lock().unwrap() = Some(( + run_ltile_primitives(None).await, + run_ltile_primitives(Some(2)).await, + )); + }); + executor.enter(Instant::ETERNITY).await; + let (fixed, affine) = result.lock().unwrap().take().unwrap(); + let rows = 2_usize; + let tiles = 8_usize; + let cols = 8_usize; + let original = (0..tiles * rows * cols) + .map(|index| 0.25 + index as f32 / 16.0) + .collect::>(); + let source = (0..tiles * cols) + .map(|index| 1.0 + index as f32 / 32.0) + .collect::>(); + let mut after_scale = vec![0_f32; tiles * rows * cols]; + for tile in 0..tiles { + for row in 0..rows { + let a = 0.5 + row as f32 / 32.0 + tile as f32 / 64.0; + for col in 0..cols { + let index = (tile * rows + row) * cols + col; + after_scale[index] = a * original[index] + 0.25 * source[tile * cols + col]; + } + } + } + let mut expected_output = Vec::with_capacity(tiles * cols); + for tile in 0..tiles { + for col in 0..cols { + expected_output.push( + (0..rows) + .map(|row| { + let scale = 0.5 + row as f32 / 32.0 + tile as f32 / 64.0; + after_scale[(tile * rows + row) * cols + col] * scale + }) + .sum::(), + ); + } + } + let mut expected_state = after_scale; + for tile in 0..tiles { + for row in 0..rows { + let scale = 0.5 + row as f32 / 32.0 + tile as f32 / 64.0; + for col in 0..cols { + expected_state[(tile * rows + row) * cols + col] += + source[tile * cols + col] * scale; + } + } + } + + assert_close(&fixed.output, &expected_output); + assert_close(&fixed.state, &expected_state); + assert_eq!(fixed.output, affine.output); + assert_eq!(fixed.state, affine.state); + assert!(fixed.matrix.bank_stall_cycles > affine.matrix.bank_stall_cycles); +} + +// A full-state tile deliberately uses a pitch that is a multiple of the bank +// count, matching the real 128-row recurrent state. Under fixed wiring, equal +// state rows from every head therefore land on the same bank words. The +// treatment adds a field-specific tile phase: two banks per 8-value state row, +// one bank per scalar row. +const FULL_ROWS: u32 = 16; +const FULL_TILES: u32 = 8; +const FULL_COLS: u32 = 8; +const FULL_SCALE_COLS: u32 = 4; +const FULL_COMPACT_SCALE_COLS: u32 = 32; +const FULL_PITCH: u32 = 16; + +#[derive(Debug)] +struct FullLTileResult { + state: Vec, + output: Vec, + matrix: MatrixPacketCounterSnapshot, + cycles: u64, +} + +fn full_region_base(region: u32) -> u32 { + region * 128 * *MLEN +} + +fn full_layout(rows: u32, cols: u32, affine: bool) -> MatrixLayout { + let words_per_row = cols / *BLEN; + MatrixLayout { + rows, + cols, + tile_count: FULL_TILES, + tile_pitch_rows: if affine { + 2 * words_per_row + } else { + FULL_PITCH + }, + alpha: 1, + tile_skew: if affine { words_per_row } else { 0 }, + } +} + +fn full_head_major_layout(rows: u32, cols: u32, affine: bool) -> MatrixLayout { + if !affine { + return full_layout(rows, cols, false); + } + MatrixLayout { + rows, + cols, + tile_count: FULL_TILES, + // Consecutive heads occupy consecutive groups of physical rows. For + // a two-row coefficient pair this makes bank=(2*head+field+word), so + // one key-column read uses each bank at most once. + tile_pitch_rows: rows, + alpha: 1, + tile_skew: 0, + } +} + +fn full_map(rows: u32, cols: u32, affine: bool, broadcast_minor: bool) -> u32 { + let _ = rows; + let words_per_row = cols / *BLEN; + ltile_map_word( + if affine { + 2 * words_per_row + } else { + FULL_PITCH + }, + affine.then_some(words_per_row), + broadcast_minor, + ) +} + +fn full_head_major_map(rows: u32, cols: u32, affine: bool) -> u32 { + if !affine { + return full_map(rows, cols, false, true); + } + let _ = cols; + ltile_map_word(rows, Some(0), true) +} + +fn append_ltile_exec( + ops: &mut Vec, + destination: (u32, u32, u32), + source: (u32, u32, u32), + scale: (u32, u32), + affine: bool, + primitive: op::LTilePrimitive, +) { + let (dst_base, dst_rows, dst_cols) = destination; + let (src_base, src_rows, src_cols) = source; + let (scale_base, scale_rows) = scale; + append_ltile_exec_with_scale_shape( + ops, + dst_base, + dst_rows, + dst_cols, + src_base, + src_rows, + src_cols, + scale_base, + scale_rows, + FULL_SCALE_COLS, + FULL_TILES, + affine, + primitive, + op::LTileAxis::Row, + ); +} + +#[allow(clippy::too_many_arguments)] +fn append_ltile_exec_with_scale_shape( + ops: &mut Vec, + dst_base: u32, + dst_rows: u32, + dst_cols: u32, + src_base: u32, + src_rows: u32, + src_cols: u32, + scale_base: u32, + scale_rows: u32, + scale_cols: u32, + scale_tiles: u32, + affine: bool, + primitive: op::LTilePrimitive, + scale_axis: op::LTileAxis, +) { + configure_ltile_view( + ops, + 0, + ltile_shape_word(dst_rows, dst_cols, FULL_TILES), + full_map(dst_rows, dst_cols, affine, false), + ); + configure_ltile_view( + ops, + 1, + ltile_shape_word(src_rows, src_cols, FULL_TILES), + full_map(src_rows, src_cols, affine, false), + ); + configure_ltile_view( + ops, + 2, + ltile_shape_word(scale_rows, scale_cols, scale_tiles), + if scale_axis == op::LTileAxis::Column { + full_head_major_map(scale_rows, scale_cols, affine) + } else { + full_map(scale_rows, scale_cols, affine, true) + }, + ); + ops.extend([ + set_gp(1, dst_base), + set_gp(2, src_base), + set_gp(3, scale_base), + op::Opcode::L_TILE_EXEC { + rd: 1, + rs1: 2, + rs2: 3, + primitive, + source_axis: op::LTileAxis::Row, + scale_axis, + }, + ]); +} + +#[allow(clippy::too_many_arguments)] +fn append_ltile_exec_compact( + ops: &mut Vec, + dst_base: u32, + dst_rows: u32, + dst_cols: u32, + src_base: u32, + src_rows: u32, + src_cols: u32, + scale_base: u32, + scale_rows: u32, + affine: bool, + primitive: op::LTilePrimitive, +) { + append_ltile_exec_with_scale_shape( + ops, + dst_base, + dst_rows, + dst_cols, + src_base, + src_rows, + src_cols, + scale_base, + scale_rows, + FULL_COMPACT_SCALE_COLS, + 1, + affine, + primitive, + op::LTileAxis::Row, + ); +} + +async fn seed_full_packet( + mram: &MatrixSram, + base: u32, + rows: u32, + cols: u32, + affine: bool, + values: &[f32], +) { + mram.write_layout_packet( + base, + full_layout(rows, cols, affine), + QuantTensor::quantize(Tensor::from_slice(values), mram.ty()), + ) + .await; +} + +async fn seed_full_head_major_packet( + mram: &MatrixSram, + base: u32, + rows: u32, + cols: u32, + affine: bool, + values: &[f32], +) { + mram.write_layout_packet( + base, + full_head_major_layout(rows, cols, affine), + QuantTensor::quantize(Tensor::from_slice(values), mram.ty()), + ) + .await; +} + +fn full_state_type() -> MxDataType { + MxDataType::Plain(DataType::Fp(FpType::BF16)) +} + +fn full_state_seed() -> Vec { + (0..FULL_TILES * FULL_ROWS * FULL_COLS) + .map(|index| 0.5 + index as f32 / 64.0) + .collect() +} + +fn round_bf16(value: f32) -> f32 { + bf16::from_f32(value).to_f32() +} + +fn full_vector_seed(offset: f32) -> Vec { + (0..FULL_TILES * FULL_COLS) + .map(|index| offset + index as f32 / 128.0) + .collect() +} + +fn full_scales(rows: u32, mut values: F) -> Vec +where + F: FnMut(u32, u32) -> (f32, f32), +{ + let mut packed = Vec::with_capacity((FULL_TILES * rows * FULL_SCALE_COLS) as usize); + for tile in 0..FULL_TILES { + for row in 0..rows { + let (a, b) = values(tile, row); + packed.extend([a, b, 0.0, 0.0]); + } + } + packed +} + +/// Keep projected per-head fields in their natural `[head][field][key]` +/// order. `L_TILE_EXEC` selects a key column and restores one scalar (or one +/// `[a,b]` pair) per head, so no copied key-major transpose is involved. +fn full_head_major_fields(field_rows: u32, mut value: F) -> Vec +where + F: FnMut(u32, u32, u32) -> f32, +{ + let mut packed = Vec::with_capacity((FULL_TILES * field_rows * FULL_ROWS) as usize); + for tile in 0..FULL_TILES { + for field in 0..field_rows { + for key in 0..FULL_ROWS { + packed.push(value(tile, field, key)); + } + } + } + packed +} + +fn full_compact_scale_pairs(rows: u32, mut values: F) -> Vec +where + F: FnMut(u32, u32) -> (f32, f32), +{ + let mut packed = Vec::with_capacity((rows * FULL_COMPACT_SCALE_COLS) as usize); + for row in 0..rows { + for tile in 0..FULL_TILES { + let (a, b) = values(tile, row); + packed.extend([a, b]); + } + packed.resize( + packed.len() + (FULL_COMPACT_SCALE_COLS - 2 * FULL_TILES) as usize, + 0.0, + ); + } + packed +} + +fn full_compact_scalars(rows: u32, mut value: F) -> Vec +where + F: FnMut(u32, u32) -> f32, +{ + let mut packed = Vec::with_capacity((rows * FULL_COMPACT_SCALE_COLS) as usize); + for row in 0..rows { + for tile in 0..FULL_TILES { + packed.push(value(tile, row)); + } + packed.resize( + packed.len() + (FULL_COMPACT_SCALE_COLS - FULL_TILES) as usize, + 0.0, + ); + } + packed +} + +async fn read_full_packet( + mram: &MatrixSram, + base: u32, + rows: u32, + cols: u32, + affine: bool, +) -> Vec { + tensor_to_f32_vec( + mram.read_layout_packet(base, full_layout(rows, cols, affine)) + .await + .0 + .as_tensor(), + ) +} + +async fn run_full_mamba_ltile(affine: bool) -> FullLTileResult { + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + let state_base = full_region_base(0); + let x_base = full_region_base(1); + let scratch_base = full_region_base(2); + let dt_base = full_region_base(3); + let update_base = full_region_base(4); + let c_base = full_region_base(5); + let output_base = full_region_base(6); + let skip_base = full_region_base(7); + let state = full_state_seed(); + let x = full_vector_seed(1.0); + let zeros = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + let dt = full_scales(1, |tile, _| (0.0, 0.25 + tile as f32 / 64.0)); + let update = full_scales(FULL_ROWS, |tile, row| { + ( + 0.75 + row as f32 / 16.0, + 0.125 + tile as f32 / 128.0 + row as f32 / 64.0, + ) + }); + let c = full_scales(FULL_ROWS, |tile, row| { + (0.25 + tile as f32 / 128.0 + row as f32 / 16.0, 0.0) + }); + let skip = full_scales(1, |tile, _| (1.0, 0.5 + tile as f32 / 128.0)); + + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 1024, + *BLEN, + full_state_type(), + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + for (base, rows, cols, values) in [ + (state_base, FULL_ROWS, FULL_COLS, state.as_slice()), + (x_base, 1, FULL_COLS, x.as_slice()), + (scratch_base, 1, FULL_COLS, zeros.as_slice()), + (dt_base, 1, FULL_SCALE_COLS, dt.as_slice()), + (update_base, FULL_ROWS, FULL_SCALE_COLS, update.as_slice()), + (c_base, FULL_ROWS, FULL_SCALE_COLS, c.as_slice()), + (output_base, 1, FULL_COLS, zeros.as_slice()), + (skip_base, 1, FULL_SCALE_COLS, skip.as_slice()), + ] { + seed_full_packet(&mram, base, rows, cols, affine, values).await; + } + + let mut ops = Vec::new(); + append_ltile_exec( + &mut ops, + (scratch_base, 1, FULL_COLS), + (x_base, 1, FULL_COLS), + (dt_base, 1), + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec( + &mut ops, + (state_base, FULL_ROWS, FULL_COLS), + (scratch_base, 1, FULL_COLS), + (update_base, FULL_ROWS), + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec( + &mut ops, + (output_base, 1, FULL_COLS), + (state_base, FULL_ROWS, FULL_COLS), + (c_base, FULL_ROWS), + affine, + op::LTilePrimitive::DotReduce, + ); + append_ltile_exec( + &mut ops, + (output_base, 1, FULL_COLS), + (x_base, 1, FULL_COLS), + (skip_base, 1), + affine, + op::LTilePrimitive::ScaleAccum, + ); + + let mut accelerator = new_accelerator(mram.clone(), vram).await; + mram.reset_packet_counters(); + let start = Executor::current().now(); + for _ in 0..TOKENS { + // State persists across tokens, but the reduction target is a + // token-local value. DOT_REDUCE reads the old destination so that + // several state chunks can accumulate into it; therefore the program + // must explicitly initialise the first chunk's accumulator. + seed_full_packet(&mram, output_base, 1, FULL_COLS, affine, &zeros).await; + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + } + let cycles = (Executor::current().now() - start).as_picos() / PERIOD.as_picos(); + let matrix = mram.packet_counter_snapshot(); + FullLTileResult { + state: read_full_packet(&mram, state_base, FULL_ROWS, FULL_COLS, affine).await, + output: read_full_packet(&mram, output_base, 1, FULL_COLS, affine).await, + matrix, + cycles, + } +} + +fn full_mamba_reference() -> (Vec, Vec) { + let mut state = full_state_seed() + .into_iter() + .map(round_bf16) + .collect::>(); + let x = full_vector_seed(1.0) + .into_iter() + .map(round_bf16) + .collect::>(); + let mut output = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + for _ in 0..TOKENS { + for tile in 0..FULL_TILES as usize { + let dt = round_bf16(0.25 + tile as f32 / 64.0); + for col in 0..FULL_COLS as usize { + let vector_index = tile * FULL_COLS as usize + col; + let scratch = round_bf16(dt * x[vector_index]); + let mut y = 0.0; + for row in 0..FULL_ROWS as usize { + let state_index = (tile * FULL_ROWS as usize + row) * FULL_COLS as usize + col; + let decay = round_bf16(0.75 + row as f32 / 16.0); + let b = round_bf16(0.125 + tile as f32 / 128.0 + row as f32 / 64.0); + state[state_index] = round_bf16(decay * state[state_index] + b * scratch); + let c = round_bf16(0.25 + tile as f32 / 128.0 + row as f32 / 16.0); + y += c * state[state_index]; + } + let reduced = round_bf16(y); + let d = round_bf16(0.5 + tile as f32 / 128.0); + output[vector_index] = round_bf16(reduced + d * x[vector_index]); + } + } + } + (output, state) +} + +async fn run_full_kda_ltile(affine: bool) -> FullLTileResult { + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + let state_base = full_region_base(0); + let zero_base = full_region_base(1); + let decay_base = full_region_base(2); + let pred_base = full_region_base(3); + let k_base = full_region_base(4); + let value_base = full_region_base(5); + let error_base = full_region_base(6); + let beta_base = full_region_base(7); + let q_base = full_region_base(8); + let output_base = full_region_base(9); + let state = full_state_seed(); + let zero = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + let value = full_vector_seed(1.5); + let decay = full_head_major_fields(2, |tile, field, key| match field { + 0 => 0.75 + tile as f32 / 128.0 + key as f32 / 32.0, + 1 => 0.0, + _ => unreachable!(), + }); + let k = full_head_major_fields(1, |tile, _, key| { + 0.125 + tile as f32 / 256.0 + key as f32 / 16.0 + }); + let beta = full_scales(1, |tile, _| { + let beta = 0.25 + tile as f32 / 128.0; + (beta, -beta) + }); + let q = full_head_major_fields(1, |tile, _, key| { + 0.25 + tile as f32 / 256.0 + key as f32 / 32.0 + }); + + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 1408, + *BLEN, + full_state_type(), + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + for (base, rows, cols, values) in [ + (state_base, FULL_ROWS, FULL_COLS, state.as_slice()), + (zero_base, 1, FULL_COLS, zero.as_slice()), + (pred_base, 1, FULL_COLS, zero.as_slice()), + (value_base, 1, FULL_COLS, value.as_slice()), + (error_base, 1, FULL_COLS, value.as_slice()), + (beta_base, 1, FULL_SCALE_COLS, beta.as_slice()), + (output_base, 1, FULL_COLS, zero.as_slice()), + ] { + seed_full_packet(&mram, base, rows, cols, affine, values).await; + } + for (base, rows, values) in [ + (decay_base, 2, decay.as_slice()), + (k_base, 1, k.as_slice()), + (q_base, 1, q.as_slice()), + ] { + seed_full_head_major_packet(&mram, base, rows, FULL_ROWS, affine, values).await; + } + + let mut ops = Vec::new(); + append_ltile_exec_with_scale_shape( + &mut ops, + state_base, + FULL_ROWS, + FULL_COLS, + zero_base, + 1, + FULL_COLS, + decay_base, + 2, + FULL_ROWS, + FULL_TILES, + affine, + op::LTilePrimitive::ScaleAccum, + op::LTileAxis::Column, + ); + append_ltile_exec_with_scale_shape( + &mut ops, + pred_base, + 1, + FULL_COLS, + state_base, + FULL_ROWS, + FULL_COLS, + k_base, + 1, + FULL_ROWS, + FULL_TILES, + affine, + op::LTilePrimitive::DotReduce, + op::LTileAxis::Column, + ); + append_ltile_exec( + &mut ops, + (error_base, 1, FULL_COLS), + (pred_base, 1, FULL_COLS), + (beta_base, 1), + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec_with_scale_shape( + &mut ops, + state_base, + FULL_ROWS, + FULL_COLS, + error_base, + 1, + FULL_COLS, + k_base, + 1, + FULL_ROWS, + FULL_TILES, + affine, + op::LTilePrimitive::OuterUpdate, + op::LTileAxis::Column, + ); + append_ltile_exec_with_scale_shape( + &mut ops, + output_base, + 1, + FULL_COLS, + state_base, + FULL_ROWS, + FULL_COLS, + q_base, + 1, + FULL_ROWS, + FULL_TILES, + affine, + op::LTilePrimitive::DotReduce, + op::LTileAxis::Column, + ); + + let mut accelerator = new_accelerator(mram.clone(), vram).await; + mram.reset_packet_counters(); + let start = Executor::current().now(); + for _ in 0..TOKENS { + // Projection produces a fresh v tensor for every token; seed the + // destination through the same affine Matrix-SRAM mapping before the + // error primitive consumes it. + seed_full_packet(&mram, error_base, 1, FULL_COLS, affine, &value).await; + // Prediction and readout reduce across one or more state chunks. They + // are token-local accumulators, unlike the persistent recurrent state. + seed_full_packet(&mram, pred_base, 1, FULL_COLS, affine, &zero).await; + seed_full_packet(&mram, output_base, 1, FULL_COLS, affine, &zero).await; + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + } + let cycles = (Executor::current().now() - start).as_picos() / PERIOD.as_picos(); + let matrix = mram.packet_counter_snapshot(); + FullLTileResult { + state: read_full_packet(&mram, state_base, FULL_ROWS, FULL_COLS, affine).await, + output: read_full_packet(&mram, output_base, 1, FULL_COLS, affine).await, + matrix, + cycles, + } +} + +fn full_kda_reference() -> (Vec, Vec) { + let mut state = full_state_seed() + .into_iter() + .map(round_bf16) + .collect::>(); + let value = full_vector_seed(1.5) + .into_iter() + .map(round_bf16) + .collect::>(); + let mut output = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + for _ in 0..TOKENS { + for tile in 0..FULL_TILES as usize { + let beta = round_bf16(0.25 + tile as f32 / 128.0); + let negative_beta = round_bf16(-beta); + for col in 0..FULL_COLS as usize { + let vector_index = tile * FULL_COLS as usize + col; + let mut prediction = 0.0; + for row in 0..FULL_ROWS as usize { + let state_index = (tile * FULL_ROWS as usize + row) * FULL_COLS as usize + col; + let decay = round_bf16(0.75 + tile as f32 / 128.0 + row as f32 / 32.0); + let k = round_bf16(0.125 + tile as f32 / 256.0 + row as f32 / 16.0); + state[state_index] = round_bf16(state[state_index] * decay); + prediction += state[state_index] * k; + } + let prediction = round_bf16(prediction); + let error = round_bf16(beta * value[vector_index] + negative_beta * prediction); + let mut readout = 0.0; + for row in 0..FULL_ROWS as usize { + let state_index = (tile * FULL_ROWS as usize + row) * FULL_COLS as usize + col; + let k = round_bf16(0.125 + tile as f32 / 256.0 + row as f32 / 16.0); + let q = round_bf16(0.25 + tile as f32 / 256.0 + row as f32 / 32.0); + state[state_index] = round_bf16(state[state_index] + error * k); + readout += state[state_index] * q; + } + output[vector_index] = round_bf16(readout); + } + } + } + (output, state) +} + +#[derive(Debug)] +struct ConnectedLTileResult { + state: Vec, + output: Vec, + matrix: MatrixPacketCounterSnapshot, + cycles: u64, +} + +#[allow(clippy::too_many_arguments)] +async fn execute_hbm_connected_program( + ops: Vec, + mram: Arc, + image: Vec, + state_output_region: u32, + state_values: usize, + output_region: u32, + output_values: usize, + scoreboard: bool, +) -> ConnectedLTileResult { + set_timing_mode(if scoreboard { + TimingMode::Scoreboard + } else { + TimingMode::Serial + }); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + let (mut accelerator, hbm) = new_accelerator_with_hbm(mram.clone(), vram, &image); + mram.reset_packet_counters(); + let start = Executor::current().now(); + if scoreboard { + let mut dependencies = Scoreboard::new(false); + accelerator + .do_ops( + &ops, + None, + TimingDriver::Scoreboard { + scoreboard: &mut dependencies, + }, + ) + .await; + } else { + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + } + let cycles = (Executor::current().now() - start).as_picos() / PERIOD.as_picos(); + ConnectedLTileResult { + state: read_hbm_state_packet(&hbm, state_output_region, state_values), + output: read_hbm_state_packet(&hbm, output_region, output_values), + matrix: mram.packet_counter_snapshot(), + cycles, + } +} + +async fn run_hbm_connected_mamba(affine: bool, scoreboard: bool) -> ConnectedLTileResult { + const STATE_IN: u32 = 0; + const X_IN: u32 = 1; + const ZERO_IN: u32 = 2; + const DT_IN: u32 = 3; + const UPDATE_IN: u32 = 4; + const C_IN: u32 = 5; + const SKIP_IN: u32 = 6; + const OUTPUT_BASE: u32 = 16; + const STATE_OUT: u32 = 31; + + let state_base = full_region_base(0); + let x_base = full_region_base(1); + let scratch_base = full_region_base(2); + let dt_base = full_region_base(3); + let update_base = full_region_base(4); + let c_base = full_region_base(5); + let output_base = full_region_base(6); + let skip_base = full_region_base(7); + let state = full_state_seed(); + let x = full_vector_seed(1.0); + let zeros = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + let dt = full_compact_scale_pairs(1, |tile, _| (0.0, 0.25 + tile as f32 / 64.0)); + let update = full_compact_scale_pairs(FULL_ROWS, |tile, row| { + ( + 0.75 + row as f32 / 16.0, + 0.125 + tile as f32 / 128.0 + row as f32 / 64.0, + ) + }); + let c = full_compact_scalars(FULL_ROWS, |tile, row| { + 0.25 + tile as f32 / 128.0 + row as f32 / 16.0 + }); + let skip = full_compact_scale_pairs(1, |tile, _| (1.0, 0.5 + tile as f32 / 128.0)); + let mut image = vec![0_u8; HBM_TEST_CAPACITY]; + for (region, values) in [ + (STATE_IN, state.as_slice()), + (X_IN, x.as_slice()), + (ZERO_IN, zeros.as_slice()), + (DT_IN, dt.as_slice()), + (UPDATE_IN, update.as_slice()), + (C_IN, c.as_slice()), + (SKIP_IN, skip.as_slice()), + ] { + write_hbm_state_packet(&mut image, region, values); + } + + let mut ops = Vec::new(); + append_matrix_view_dma( + &mut ops, true, state_base, STATE_IN, FULL_ROWS, FULL_COLS, affine, false, + ); + for token in 0..TOKENS as u32 { + for (base, region, rows, cols, broadcast) in [ + (x_base, X_IN, 1, FULL_COLS, false), + (scratch_base, ZERO_IN, 1, FULL_COLS, false), + (dt_base, DT_IN, 1, FULL_COMPACT_SCALE_COLS, true), + ( + update_base, + UPDATE_IN, + FULL_ROWS, + FULL_COMPACT_SCALE_COLS, + true, + ), + (c_base, C_IN, FULL_ROWS, FULL_COMPACT_SCALE_COLS, true), + (output_base, ZERO_IN, 1, FULL_COLS, false), + (skip_base, SKIP_IN, 1, FULL_COMPACT_SCALE_COLS, true), + ] { + append_matrix_view_dma(&mut ops, true, base, region, rows, cols, affine, broadcast); + } + append_ltile_exec_compact( + &mut ops, + scratch_base, + 1, + FULL_COLS, + x_base, + 1, + FULL_COLS, + dt_base, + 1, + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec_compact( + &mut ops, + state_base, + FULL_ROWS, + FULL_COLS, + scratch_base, + 1, + FULL_COLS, + update_base, + FULL_ROWS, + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec_compact( + &mut ops, + output_base, + 1, + FULL_COLS, + state_base, + FULL_ROWS, + FULL_COLS, + c_base, + FULL_ROWS, + affine, + op::LTilePrimitive::DotReduce, + ); + append_ltile_exec_compact( + &mut ops, + output_base, + 1, + FULL_COLS, + x_base, + 1, + FULL_COLS, + skip_base, + 1, + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_matrix_view_dma( + &mut ops, + false, + output_base, + OUTPUT_BASE + token, + 1, + FULL_COLS, + affine, + false, + ); + } + append_matrix_view_dma( + &mut ops, false, state_base, STATE_OUT, FULL_ROWS, FULL_COLS, affine, false, + ); + + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 1024, + *BLEN, + full_state_type(), + )); + execute_hbm_connected_program( + ops, + mram, + image, + STATE_OUT, + state.len(), + OUTPUT_BASE + TOKENS as u32 - 1, + zeros.len(), + scoreboard, + ) + .await +} + +async fn run_hbm_connected_kda(affine: bool, scoreboard: bool) -> ConnectedLTileResult { + const STATE_IN: u32 = 0; + const ZERO_IN: u32 = 1; + const DECAY_IN: u32 = 2; + const K_IN: u32 = 3; + const VALUE_IN: u32 = 4; + const BETA_IN: u32 = 5; + const Q_IN: u32 = 6; + const OUTPUT_BASE: u32 = 16; + const STATE_OUT: u32 = 31; + + let state_base = full_region_base(0); + let zero_base = full_region_base(1); + let decay_base = full_region_base(2); + let pred_base = full_region_base(3); + let k_base = full_region_base(4); + let value_base = full_region_base(5); + let error_base = full_region_base(6); + let beta_base = full_region_base(7); + let q_base = full_region_base(8); + let output_base = full_region_base(9); + let state = full_state_seed(); + let zero = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + let value = full_vector_seed(1.5); + let decay = full_compact_scale_pairs(FULL_ROWS, |tile, row| { + (0.75 + tile as f32 / 128.0 + row as f32 / 32.0, 0.0) + }); + let k = full_compact_scalars(FULL_ROWS, |tile, row| { + 0.125 + tile as f32 / 256.0 + row as f32 / 16.0 + }); + let beta = full_compact_scale_pairs(1, |tile, _| { + let beta = 0.25 + tile as f32 / 128.0; + (beta, -beta) + }); + let q = full_compact_scalars(FULL_ROWS, |tile, row| { + 0.25 + tile as f32 / 256.0 + row as f32 / 32.0 + }); + let mut image = vec![0_u8; HBM_TEST_CAPACITY]; + for (region, values) in [ + (STATE_IN, state.as_slice()), + (ZERO_IN, zero.as_slice()), + (DECAY_IN, decay.as_slice()), + (K_IN, k.as_slice()), + (VALUE_IN, value.as_slice()), + (BETA_IN, beta.as_slice()), + (Q_IN, q.as_slice()), + ] { + write_hbm_state_packet(&mut image, region, values); + } + + let mut ops = Vec::new(); + append_matrix_view_dma( + &mut ops, true, state_base, STATE_IN, FULL_ROWS, FULL_COLS, affine, false, + ); + append_matrix_view_dma( + &mut ops, true, zero_base, ZERO_IN, 1, FULL_COLS, affine, false, + ); + for token in 0..TOKENS as u32 { + for (base, region, rows, cols, broadcast) in [ + ( + decay_base, + DECAY_IN, + FULL_ROWS, + FULL_COMPACT_SCALE_COLS, + true, + ), + (pred_base, ZERO_IN, 1, FULL_COLS, false), + (k_base, K_IN, FULL_ROWS, FULL_COMPACT_SCALE_COLS, true), + (value_base, VALUE_IN, 1, FULL_COLS, false), + (error_base, VALUE_IN, 1, FULL_COLS, false), + (beta_base, BETA_IN, 1, FULL_COMPACT_SCALE_COLS, true), + (q_base, Q_IN, FULL_ROWS, FULL_COMPACT_SCALE_COLS, true), + (output_base, ZERO_IN, 1, FULL_COLS, false), + ] { + append_matrix_view_dma(&mut ops, true, base, region, rows, cols, affine, broadcast); + } + append_ltile_exec_compact( + &mut ops, + state_base, + FULL_ROWS, + FULL_COLS, + zero_base, + 1, + FULL_COLS, + decay_base, + FULL_ROWS, + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec_compact( + &mut ops, + pred_base, + 1, + FULL_COLS, + state_base, + FULL_ROWS, + FULL_COLS, + k_base, + FULL_ROWS, + affine, + op::LTilePrimitive::DotReduce, + ); + append_ltile_exec_compact( + &mut ops, + error_base, + 1, + FULL_COLS, + pred_base, + 1, + FULL_COLS, + beta_base, + 1, + affine, + op::LTilePrimitive::ScaleAccum, + ); + append_ltile_exec_compact( + &mut ops, + state_base, + FULL_ROWS, + FULL_COLS, + error_base, + 1, + FULL_COLS, + k_base, + FULL_ROWS, + affine, + op::LTilePrimitive::OuterUpdate, + ); + append_ltile_exec_compact( + &mut ops, + output_base, + 1, + FULL_COLS, + state_base, + FULL_ROWS, + FULL_COLS, + q_base, + FULL_ROWS, + affine, + op::LTilePrimitive::DotReduce, + ); + append_matrix_view_dma( + &mut ops, + false, + output_base, + OUTPUT_BASE + token, + 1, + FULL_COLS, + affine, + false, + ); + } + append_matrix_view_dma( + &mut ops, false, state_base, STATE_OUT, FULL_ROWS, FULL_COLS, affine, false, + ); + + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 1408, + *BLEN, + full_state_type(), + )); + execute_hbm_connected_program( + ops, + mram, + image, + STATE_OUT, + state.len(), + OUTPUT_BASE + TOKENS as u32 - 1, + zero.len(), + scoreboard, + ) + .await +} + +#[tokio::test] +async fn dot_reduce_accumulates_across_two_state_chunks() { + const CHUNK_ROWS: u32 = FULL_ROWS / 2; + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + let first_state_base = full_region_base(0); + let second_state_base = full_region_base(1); + let first_scale_base = full_region_base(2); + let second_scale_base = full_region_base(3); + let output_base = full_region_base(4); + let first = (0..FULL_TILES * CHUNK_ROWS * FULL_COLS) + .map(|index| 0.25 + index as f32 / 64.0) + .collect::>(); + let second = (0..FULL_TILES * CHUNK_ROWS * FULL_COLS) + .map(|index| 1.0 + index as f32 / 32.0) + .collect::>(); + let first_scale = full_scales(CHUNK_ROWS, |tile, row| { + (0.5 + tile as f32 / 128.0 + row as f32 / 64.0, 0.0) + }); + let second_scale = full_scales(CHUNK_ROWS, |tile, row| { + (0.75 + tile as f32 / 128.0 + row as f32 / 32.0, 0.0) + }); + let zero = vec![0.0; (FULL_TILES * FULL_COLS) as usize]; + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MLEN as usize * 640, + *BLEN, + full_state_type(), + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 32, *VECTOR_SRAM_TYPE)); + for (base, rows, cols, values) in [ + (first_state_base, CHUNK_ROWS, FULL_COLS, first.as_slice()), + (second_state_base, CHUNK_ROWS, FULL_COLS, second.as_slice()), + ( + first_scale_base, + CHUNK_ROWS, + FULL_SCALE_COLS, + first_scale.as_slice(), + ), + ( + second_scale_base, + CHUNK_ROWS, + FULL_SCALE_COLS, + second_scale.as_slice(), + ), + (output_base, 1, FULL_COLS, zero.as_slice()), + ] { + seed_full_packet(&mram, base, rows, cols, true, values).await; + } + let mut ops = Vec::new(); + for (state_base, scale_base) in [ + (first_state_base, first_scale_base), + (second_state_base, second_scale_base), + ] { + append_ltile_exec( + &mut ops, + (output_base, 1, FULL_COLS), + (state_base, CHUNK_ROWS, FULL_COLS), + (scale_base, CHUNK_ROWS), + true, + op::LTilePrimitive::DotReduce, + ); + } + let mut accelerator = new_accelerator(mram.clone(), vram).await; + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + let output = read_full_packet(&mram, output_base, 1, FULL_COLS, true).await; + let mut expected = vec![0.0; output.len()]; + for tile in 0..FULL_TILES as usize { + for col in 0..FULL_COLS as usize { + let output_index = tile * FULL_COLS as usize + col; + for row in 0..CHUNK_ROWS as usize { + let index = (tile * CHUNK_ROWS as usize + row) * FULL_COLS as usize + col; + expected[output_index] += + first[index] * (0.5 + tile as f32 / 128.0 + row as f32 / 64.0); + expected[output_index] += + second[index] * (0.75 + tile as f32 / 128.0 + row as f32 / 32.0); + } + } + } + *task_result.lock().unwrap() = Some((output, expected)); + }); + executor.enter(Instant::ETERNITY).await; + let (output, expected) = result.lock().unwrap().take().unwrap(); + assert_close(&output, &expected); +} + +#[tokio::test] +async fn full_mamba_recurrence_uses_four_ltile_execs_across_four_tokens() { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + *task_result.lock().unwrap() = Some(( + run_full_mamba_ltile(false).await, + run_full_mamba_ltile(true).await, + )); + }); + executor.enter(Instant::ETERNITY).await; + let (fixed, affine) = result.lock().unwrap().take().unwrap(); + eprintln!( + "Mamba full recurrence: fixed cycles={} stalls={}, affine cycles={} stalls={}", + fixed.cycles, + fixed.matrix.bank_stall_cycles, + affine.cycles, + affine.matrix.bank_stall_cycles, + ); + let (expected_output, expected_state) = full_mamba_reference(); + assert_close(&fixed.output, &expected_output); + assert_close(&fixed.state, &expected_state); + assert_eq!(fixed.output, affine.output); + assert_eq!(fixed.state, affine.state); + assert!(fixed.matrix.bank_stall_cycles > 0); + assert_eq!(affine.matrix.bank_stall_cycles, 0); + assert!(fixed.cycles > affine.cycles); +} + +#[tokio::test] +async fn full_kda_recurrence_uses_five_ltile_execs_across_four_tokens() { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + *task_result.lock().unwrap() = Some(( + run_full_kda_ltile(false).await, + run_full_kda_ltile(true).await, + )); + }); + executor.enter(Instant::ETERNITY).await; + let (fixed, affine) = result.lock().unwrap().take().unwrap(); + eprintln!( + "KDA full recurrence: fixed cycles={} stalls={}, affine cycles={} stalls={}", + fixed.cycles, + fixed.matrix.bank_stall_cycles, + affine.cycles, + affine.matrix.bank_stall_cycles, + ); + let (expected_output, expected_state) = full_kda_reference(); + assert_close(&fixed.output, &expected_output); + assert_close(&fixed.state, &expected_state); + assert_eq!(fixed.output, affine.output); + assert_eq!(fixed.state, affine.state); + assert!(fixed.matrix.bank_stall_cycles > 0); + assert_eq!(affine.matrix.bank_stall_cycles, 0); + assert!(fixed.cycles > affine.cycles); +} + +#[tokio::test] +async fn mamba_hbm_to_affine_matrix_to_hbm_is_numerically_connected() { + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + *task_result.lock().unwrap() = Some(( + run_hbm_connected_mamba(false, false).await, + run_hbm_connected_mamba(true, false).await, + run_hbm_connected_mamba(true, true).await, + )); + }); + executor.enter(Instant::ETERNITY).await; + let (fixed, affine, pipelined) = result.lock().unwrap().take().unwrap(); + let (expected_output, expected_state) = full_mamba_reference(); + + assert_close(&fixed.output, &expected_output); + assert_close(&fixed.state, &expected_state); + assert_eq!(fixed.output, affine.output); + assert_eq!(fixed.state, affine.state); + assert_eq!(affine.output, pipelined.output); + assert_eq!(affine.state, pipelined.state); + assert!(fixed.matrix.bank_stall_cycles > 0); + assert_eq!(affine.matrix.bank_stall_cycles, 0); + assert!(fixed.cycles > affine.cycles); +} + +#[tokio::test] +async fn kda_hbm_to_affine_matrix_to_hbm_is_numerically_connected() { + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let task_result = result.clone(); + executor.spawn(async move { + *task_result.lock().unwrap() = Some(( + run_hbm_connected_kda(false, false).await, + run_hbm_connected_kda(true, false).await, + run_hbm_connected_kda(true, true).await, + )); + }); + executor.enter(Instant::ETERNITY).await; + let (fixed, affine, pipelined) = result.lock().unwrap().take().unwrap(); + let (expected_output, expected_state) = full_kda_reference(); + + assert_close(&fixed.output, &expected_output); + assert_close(&fixed.state, &expected_state); + assert_eq!(fixed.output, affine.output); + assert_eq!(fixed.state, affine.state); + assert_eq!(affine.output, pipelined.output); + assert_eq!(affine.state, pipelined.state); + assert!(fixed.matrix.bank_stall_cycles > 0); + assert_eq!(affine.matrix.bank_stall_cycles, 0); + assert!(fixed.cycles > affine.cycles); +} + +async fn run_matrix_view_packet_roundtrip(tile_pitch_rows: u32) -> (Vec, u64, u64) { + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + (*MLEN as usize) * 64, + *BLEN, + *MATRIX_SRAM_TYPE, + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 8, *VECTOR_SRAM_TYPE)); + let values = (0..*VLEN) + .map(|value| value as f32 + 1.0) + .collect::>(); + let layout = sram::matrix::MatrixLayout { + rows: 1, + cols: 2 * *BLEN, + tile_count: *VLEN / (2 * *BLEN), + tile_pitch_rows, + alpha: 1, + tile_skew: 0, + }; + mram.write_layout_packet( + 0, + layout, + QuantTensor::quantize(Tensor::from_slice(&values), *MATRIX_SRAM_TYPE), + ) + .await; + mram.reset_packet_counters(); + + let m_machine = MatrixMachine::new( + mram.clone(), + vram.clone(), + *MLEN, + *HLEN, + *BLEN, + *BROADCAST_AMOUNT, + ); + let v_machine = VectorMachine::new(vram.clone(), *VLEN, *HLEN); + let hbm: Arc = Arc::new(WithTiming::new( + NaiveTiming::preset_ddr4_2400p(4), + MemoryBacked::with_capacity(4096), + )); + let mut accelerator = Accelerator::new(m_machine, v_machine, hbm); + // Eight 1x8 tiles form one 64-value packet. Pitch 1 overlaps adjacent + // two-word tiles on the fixed diagonal banks; pitch 2 uses every bank once. + let shape = (7 << 12) | (7 << 24); + let mapping = tile_pitch_rows; + let ops = vec![ + set_gp(1, shape), + set_gp(2, mapping), + set_gp(3, *VLEN), + set_gp(4, 0), + set_gp(5, 0), + op::Opcode::L_TILE_CFG { + shape: 1, + mapping: 2, + slot: 1, + }, + op::Opcode::V_ADD_VV { + rd: 3, + rs1: 4, + rs2: 5, + rmask: 0, + // Explicit Matrix marker + source-1 slot. + view_mask: 0b010, + }, + ]; + let start = Executor::current().now(); + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + let cycles = (Executor::current().now() - start).as_picos() / PERIOD.as_picos(); + let counters = accelerator.matrix_view_packet_counters(); + ( + tensor_to_f32_vec(vram.read(*VLEN).await.as_tensor()), + counters.bank_stall_cycles, + cycles, + ) +} + +#[tokio::test] +async fn l_mview_dispatch_roundtrips_values_and_removes_real_matrix_bank_conflicts() { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let result_task = result.clone(); + executor.spawn(async move { + let pitch_one = run_matrix_view_packet_roundtrip(1).await; + let co_layout = run_matrix_view_packet_roundtrip(2).await; + *result_task.lock().unwrap() = Some((pitch_one, co_layout)); + }); + executor.enter(Instant::ETERNITY).await; + let ((row_values, row_stalls, row_cycles), (affine_values, affine_stalls, affine_cycles)) = + result.lock().unwrap().take().unwrap(); + assert_eq!(row_values, affine_values); + assert!(row_stalls > 0); + assert_eq!(affine_stalls, 0); + assert!(row_cycles > affine_cycles); +} + +async fn run_matrix_accumulator_view_writeback(tile_pitch_rows: u32) -> (Vec, u64, u64) { + assert_eq!((*MLEN, *BLEN, *VLEN), (64, 4, 64)); + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + (*MLEN as usize) * 64, + *BLEN, + *MATRIX_SRAM_TYPE, + )); + let vram = Arc::new(VectorSram::from_mx_type(*VLEN, 128, *VECTOR_SRAM_TYPE)); + + let mut identity = vec![0.0_f32; (*MLEN * *MLEN) as usize]; + for index in 0..*MLEN as usize { + identity[index * *MLEN as usize + index] = 1.0; + } + mram.write( + 0, + QuantTensor::quantize(Tensor::from_slice(&identity), *MATRIX_SRAM_TYPE), + ) + .await; + + let output_blocks = *VLEN / *BLEN; + for tile in 0..output_blocks { + for row in 0..*BLEN { + let mut values = vec![0.0_f32; *VLEN as usize]; + if row == 0 { + for column in 0..*BLEN { + values[column as usize] = (tile * *BLEN + column + 1) as f32; + } + } + vram.write( + (tile * *BLEN + row) * *VLEN, + QuantTensor::quantize(Tensor::from_slice(&values), *VECTOR_SRAM_TYPE), + ) + .await; + } + } + + let m_machine = MatrixMachine::new(mram, vram.clone(), *MLEN, *HLEN, *BLEN, *BROADCAST_AMOUNT); + let v_machine = VectorMachine::new(vram.clone(), *VLEN, *HLEN); + let hbm: Arc = Arc::new(WithTiming::new( + NaiveTiming::preset_ddr4_2400p(4), + MemoryBacked::with_capacity(4096), + )); + let mut accelerator = Accelerator::new(m_machine, v_machine, hbm); + + // The producer writes sixteen BLEN-wide fragments. The consumer sees eight + // logical heads with two fragments per head, so this catches the old but + // incorrect 16x4 producer-only descriptor. + let consumer_cols = 2 * *BLEN; + let consumer_tiles = *VLEN / consumer_cols; + let shape = (consumer_cols - 1) << 12 | ((consumer_tiles - 1) << 24); + let mapping = tile_pitch_rows; + let matrix_output_base = *MLEN * *MLEN; + let vector_output_base = output_blocks * *BLEN * *VLEN; + let mut ops = Vec::new(); + ops.extend([ + set_gp(1, shape), + set_gp(2, mapping), + set_gp(3, 0), + set_gp(4, matrix_output_base), + set_gp(5, vector_output_base), + set_gp(8, vector_output_base + *VLEN), + op::Opcode::L_TILE_CFG { + shape: 1, + mapping: 2, + slot: 0, + }, + op::Opcode::L_TILE_CFG { + shape: 1, + mapping: 2, + slot: 1, + }, + ]); + for tile in 0..output_blocks { + ops.extend([ + set_gp(6, tile * *BLEN * *VLEN), + set_gp(7, tile * *BLEN), + op::Opcode::M_MM { + rs1: 3, + rs2: 6, + view: None, + }, + op::Opcode::M_MM_WO { + rd: 4, + rstride: 7, + imm: 0, + view: Some(0), + }, + ]); + } + ops.push(op::Opcode::V_ADD_VV { + rd: 5, + rs1: 4, + rs2: 8, + rmask: 0, + view_mask: 0b010, + }); + + let start = Executor::current().now(); + accelerator.do_ops(&ops, None, TimingDriver::Serial).await; + let cycles = (Executor::current().now() - start).as_picos() / PERIOD.as_picos(); + let counters = accelerator.matrix_view_packet_counters(); + ( + tensor_to_f32_vec(vram.read(vector_output_base).await.as_tensor()), + counters.bank_stall_cycles, + cycles, + ) +} + +#[tokio::test] +async fn matrix_accumulator_writes_skewed_tiles_consumed_without_bank_conflicts() { + set_timing_mode(TimingMode::Serial); + let executor = Executor::new(); + let result = Arc::new(Mutex::new(None)); + let result_task = result.clone(); + executor.spawn(async move { + let fixed = run_matrix_accumulator_view_writeback(1).await; + let affine = run_matrix_accumulator_view_writeback(2).await; + *result_task.lock().unwrap() = Some((fixed, affine)); + }); + executor.enter(Instant::ETERNITY).await; + + let ((fixed_values, fixed_stalls, fixed_cycles), (affine_values, affine_stalls, affine_cycles)) = + result.lock().unwrap().take().unwrap(); + let expected = (1..=*VLEN).map(|value| value as f32).collect::>(); + assert_eq!(fixed_values, expected); + assert_eq!(affine_values, expected); + assert!(fixed_stalls > 0); + assert_eq!(affine_stalls, 0); + assert!(fixed_cycles > affine_cycles); +} diff --git a/transactional_emulator/src/accelerator/mod.rs b/transactional_emulator/src/accelerator/mod.rs index 649598fd..6bd868a0 100644 --- a/transactional_emulator/src/accelerator/mod.rs +++ b/transactional_emulator/src/accelerator/mod.rs @@ -8,6 +8,7 @@ use std::sync::Arc; use memory::ErasedMemoryModel; +use sram::matrix::MatrixPacketCounterSnapshot; use crate::matrix_machine::MatrixMachine; use crate::vector_machine::VectorMachine; @@ -15,6 +16,9 @@ use crate::vector_machine::VectorMachine; mod access; mod dispatch; mod loop_state; +mod matrix_view; +#[cfg(test)] +mod matrix_view_tests; #[cfg(test)] mod pipeline_tests; mod registers; @@ -23,6 +27,7 @@ mod scoreboard; pub(crate) use access::Unit; pub(crate) use dispatch::TimingDriver; +pub(crate) use matrix_view::MatrixViewDescriptor; pub(crate) use scoreboard::Scoreboard; use loop_state::LoopState; @@ -44,11 +49,13 @@ impl Accelerator { v_machine: VectorMachine, hbm: Arc, ) -> Self { + let mview_banks = m_machine.mram.banks(); + let mview_bank_width = m_machine.mram.bank_width(); Self { m_machine, v_machine, hbm, - reg_file: AcceleratorRegFile::new(), + reg_file: AcceleratorRegFile::new_with_matrix(mview_banks, mview_bank_width), scalar_sram: ScalarSram::new(), loop_state: LoopState::new(), } @@ -95,4 +102,8 @@ impl Accelerator { pub(crate) fn intsram_dump_bytes(&self) -> Vec { self.scalar_sram.intsram_to_le_bytes() } + + pub(crate) fn matrix_view_packet_counters(&self) -> MatrixPacketCounterSnapshot { + self.m_machine.mram.packet_counter_snapshot() + } } diff --git a/transactional_emulator/src/accelerator/pipeline_tests.rs b/transactional_emulator/src/accelerator/pipeline_tests.rs index f100cf53..2507f189 100644 --- a/transactional_emulator/src/accelerator/pipeline_tests.rs +++ b/transactional_emulator/src/accelerator/pipeline_tests.rs @@ -122,7 +122,11 @@ fn independent_scalar_op() -> op::Opcode { } fn matrix_plus_scalars(scalars: usize) -> Vec { - let mut ops = vec![op::Opcode::M_MM { rs1: 1, rs2: 2 }]; + let mut ops = vec![op::Opcode::M_MM { + rs1: 1, + rs2: 2, + view: None, + }]; for _ in 0..scalars { ops.push(independent_scalar_op()); } @@ -188,6 +192,7 @@ async fn serialize_mode_reproduces_serial_cycle_counts() { rs1: 5, rs2: 5, rmask: 0, + view_mask: 0, }); ops }; @@ -244,6 +249,7 @@ fn prefetch_program(independent: usize, dependent: bool) -> Vec { rs1: 5, rs2: 5, rmask: 0, + view_mask: 0, }); } if dependent { @@ -253,6 +259,7 @@ fn prefetch_program(independent: usize, dependent: bool) -> Vec { rs1: 4, rs2: 4, rmask: 0, + view_mask: 0, }); } ops @@ -354,6 +361,7 @@ fn store_program(independent: usize) -> Vec { rs1: 5, rs2: 5, rmask: 0, + view_mask: 0, }, ]; for _ in 0..independent { @@ -362,6 +370,7 @@ fn store_program(independent: usize) -> Vec { rs1: 5, rs2: 5, rmask: 0, + view_mask: 0, }); } ops @@ -405,12 +414,21 @@ async fn async_store_overlaps_and_snapshots_against_war() { #[tokio::test] async fn back_to_back_matrix_ops_serialize_on_the_matrix_unit() { let ops = vec![ - op::Opcode::M_MM { rs1: 1, rs2: 2 }, - op::Opcode::M_MM { rs1: 1, rs2: 2 }, + op::Opcode::M_MM { + rs1: 1, + rs2: 2, + view: None, + }, + op::Opcode::M_MM { + rs1: 1, + rs2: 2, + view: None, + }, op::Opcode::M_MM_WO { rd: 1, rstride: 0, imm: 0, + view: None, }, ]; let pipelined = run_program(ops, RunMode::Scoreboard, vec![]).await; diff --git a/transactional_emulator/src/accelerator/registers.rs b/transactional_emulator/src/accelerator/registers.rs index 0ce66400..d78d7eac 100644 --- a/transactional_emulator/src/accelerator/registers.rs +++ b/transactional_emulator/src/accelerator/registers.rs @@ -2,6 +2,8 @@ use half::bf16; +use super::matrix_view::{MatrixViewDescriptor, MatrixViewTable}; + pub(super) struct AcceleratorRegFile { // === ISA-indexed register banks === gp_reg: [u32; 16], @@ -21,10 +23,16 @@ pub(super) struct AcceleratorRegFile { /// `topk > 0` — the program would abort with "topk must be positive", which /// says nothing about the missing `C_SET_TOPK_REG`. topk_policy: Option, + mviews: MatrixViewTable, } impl AcceleratorRegFile { + #[cfg(test)] pub(super) fn new() -> Self { + Self::new_with_matrix(16, 1) + } + + pub(super) fn new_with_matrix(banks: u32, bank_width: u32) -> Self { Self { gp_reg: [0; 16], fp_reg: [bf16::ZERO; 8], @@ -36,6 +44,7 @@ impl AcceleratorRegFile { bmm_scale: 0.25, v_mask: 0, topk_policy: None, + mviews: MatrixViewTable::new(banks, bank_width), } } @@ -118,6 +127,23 @@ impl AcceleratorRegFile { .map(|packed| ((packed >> 8) as usize, (packed & 0xFF) as usize)) } + pub(super) fn configure_mview( + &mut self, + slot: u8, + shape_register: u8, + map_register: u8, + ) -> Result<(), String> { + self.mviews.configure( + slot, + self.read_gp(shape_register), + self.read_gp(map_register), + ) + } + + pub(super) fn matrix_view(&self, slot: u8) -> Result { + self.mviews.get(slot) + } + /// `dst_gp = op(read_gp(src1), read_gp(src2))`. Helper for binary GP-to-GP /// instructions (S_ADD_INT / S_SUB_INT / S_MUL_INT). pub(super) fn binop_gp u32>( diff --git a/transactional_emulator/src/cli.rs b/transactional_emulator/src/cli.rs index 4f2de465..0294f5e7 100644 --- a/transactional_emulator/src/cli.rs +++ b/transactional_emulator/src/cli.rs @@ -165,6 +165,11 @@ pub(crate) struct Opts { /// the default ../plena_settings.toml lookup. pub(crate) settings: Option, + #[arg(long, help_heading = "Diagnostics")] + /// Optional path for the post-run HBM image. Unlike DEBUG logging this is + /// explicit and therefore suitable for numerical integration tests. + pub(crate) hbm_dump: Option, + #[arg(long)] /// Optional generated ASM source used to derive PC-to-stage labels for a /// runtime stage profile. This is diagnostic only; normal runs omit it. diff --git a/transactional_emulator/src/dma.rs b/transactional_emulator/src/dma.rs index a3222e14..f042efc5 100644 --- a/transactional_emulator/src/dma.rs +++ b/transactional_emulator/src/dma.rs @@ -18,7 +18,7 @@ use std::sync::Arc; use memory::ErasedMemoryModel; -use quantize::{DataType, MxDataType, QuantTensor, tensor_from_f32_slice}; +use quantize::{DataType, MxDataType, QuantTensor, tensor_from_f32_slice, tensor_to_f32_vec}; use runtime::Executor; use sram::VectorSram; use tokio::sync::oneshot::{self, Receiver}; @@ -88,6 +88,17 @@ impl MxLayout { scale_len_in_bytes, } } + + fn contiguous_stride_bytes(hbm_type: MxDataType, dim: u32) -> u32 { + let bits = u32::from(hbm_type.element_type().size_in_bits()) + .checked_mul(dim) + .expect("contiguous HBM row size overflowed u32"); + assert!( + bits.is_multiple_of(8), + "a contiguous HBM row must occupy a whole number of bytes" + ); + bits / 8 + } } /// A strided MX-format region in HBM — the "where + what" of a transfer, @@ -144,7 +155,14 @@ pub(crate) fn transfer_mx_from_hbm( rstride, stride, } = region; - let stride = if rstride == 1 { stride } else { load_dim }; + // HBM addresses and C_SET_STRIDE_REG are byte based. The old default + // happened to be correct for 8-bit tensors, but overlapped adjacent rows + // for BF16/FP32 because it advanced by an element count instead of bytes. + let stride = if rstride == 1 { + stride + } else { + MxLayout::contiguous_stride_bytes(hbm_type, load_dim) + }; let hbm = hbm.clone(); Executor::current().spawn(async move { @@ -326,6 +344,28 @@ pub(crate) async fn snapshot_vram_rows( rows } +/// Split one logical Matrix-view packet into contiguous HBM transfer rows. +/// +/// Matrix views restore logical tile/row/lane order before this helper runs, +/// so the HBM image stays compiler-defined packet-major and does not encode +/// the SRAM bank mapping. +pub(crate) fn split_packet_rows( + packet: &QuantTensor, + row_dim: u32, + sram_type: MxDataType, +) -> Vec { + assert!(row_dim > 0, "DMA row width must be positive"); + let values = tensor_to_f32_vec(packet.as_tensor()); + values + .chunks(row_dim as usize) + .map(|row| { + let mut padded = vec![0_f32; row_dim as usize]; + padded[..row.len()].copy_from_slice(row); + QuantTensor::quantize(tensor_from_f32_slice(&padded), sram_type) + }) + .collect() +} + /// Write the snapshotted rows into an HBM [`MxRegion`] with a strided writing /// pattern (the timed half of `H_STORE_V`). pub(crate) async fn store_rows_to_hbm( @@ -341,7 +381,11 @@ pub(crate) async fn store_rows_to_hbm( rstride, stride, } = region; - let stride = if rstride == 1 { stride } else { store_dim }; + let stride = if rstride == 1 { + stride + } else { + MxLayout::contiguous_stride_bytes(hbm_type, store_dim) + }; let layout = MxLayout::compute(hbm_type, stride, store_dim); let len_in_bytes_per_store = layout.len_in_bytes; @@ -522,4 +566,26 @@ mod tests { assert_eq!(layout.element_bits, 16); assert_eq!(layout.len_in_bytes, 128); // 16 * 64 / 8 } + + #[test] + fn test_contiguous_stride_is_measured_in_bytes_for_wide_elements() { + let bf16 = MxDataType::Plain(DataType::Fp(FpType::BF16)); + let fp32 = MxDataType::Plain(DataType::Fp(FpType::F32)); + assert_eq!(MxLayout::contiguous_stride_bytes(bf16, 64), 128); + assert_eq!(MxLayout::contiguous_stride_bytes(fp32, 64), 256); + } + + #[test] + fn test_split_packet_rows_pads_only_the_final_physical_row() { + let fp32 = MxDataType::Plain(DataType::Fp(FpType::F32)); + let values = (0..70).map(|value| value as f32).collect::>(); + let packet = QuantTensor::quantize(tensor_from_f32_slice(&values), fp32); + let rows = split_packet_rows(&packet, 64, fp32); + assert_eq!(rows.len(), 2); + let first = tensor_to_f32_vec(rows[0].as_tensor()); + let second = tensor_to_f32_vec(rows[1].as_tensor()); + assert_eq!(first, values[..64]); + assert_eq!(&second[..6], &values[64..]); + assert!(second[6..].iter().all(|value| *value == 0.0)); + } } diff --git a/transactional_emulator/src/load_config.rs b/transactional_emulator/src/load_config.rs index 5b73badb..c27635e4 100644 --- a/transactional_emulator/src/load_config.rs +++ b/transactional_emulator/src/load_config.rs @@ -124,12 +124,28 @@ pub struct PrecisionSection { pub hbm_v_act_type: MxDataTypeConfig, #[serde(rename = "HBM_V_KV_TYPE")] pub hbm_v_kv_type: MxDataTypeConfig, + /// State transfers through explicit Matrix-view DMA use their own format. + #[serde(rename = "HBM_STATE_TYPE", default = "default_state_type")] + pub hbm_state_type: MxDataTypeConfig, #[serde(rename = "HBM_V_INT_TYPE")] pub hbm_v_int_type: MxDataTypeConfig, #[serde(rename = "SCALAR_FP")] pub scalar_fp: DataTypeConfig, } +fn default_state_type() -> MxDataTypeConfig { + MxDataTypeConfig { + format: "Plain".to_string(), + data: MxDataTypeData::Plain { + data_type: DataTypeConfig::Fp(FpTypeConfig { + sign: true, + exponent: 8, + mantissa: 7, + }), + }, + } +} + #[derive(Debug, Serialize, Deserialize, Clone)] pub struct LatencySection { #[serde(rename = "SYSTOLIC_PROCESSING_OVERHEAD")] @@ -267,6 +283,7 @@ impl Default for AcceleratorConfig { }), }, }, + hbm_state_type: default_state_type(), hbm_v_int_type: MxDataTypeConfig { format: "Plain".to_string(), data: MxDataTypeData::Plain { @@ -487,6 +504,10 @@ pub fn vector_kv_type() -> MxDataType { CONFIG.precision.hbm_v_kv_type.clone().into() } +pub fn state_type() -> MxDataType { + CONFIG.precision.hbm_state_type.clone().into() +} + /// Reserved for future scalar FP ops; not yet wired into any opcode dispatch. #[allow(dead_code)] pub fn scalar_fp_type() -> DataType { diff --git a/transactional_emulator/src/matrix_machine.rs b/transactional_emulator/src/matrix_machine.rs index ce10045e..4e3a394b 100644 --- a/transactional_emulator/src/matrix_machine.rs +++ b/transactional_emulator/src/matrix_machine.rs @@ -15,9 +15,11 @@ use std::sync::Arc; use quantize::QuantTensor; +use sram::matrix::MatrixAccessAxis; use sram::{MatrixSram, VectorSram, assert_multiple_of, multiple_and_offset}; use tch::{IndexOp, Tensor}; +use crate::accelerator::MatrixViewDescriptor; use crate::matrix_core::{MatrixCore, MatrixCoreProfile}; use crate::runtime_config::SYSTOLIC_PROCESSING_OVERHEAD; @@ -97,11 +99,53 @@ impl MatrixMachine { self.core.profile() } + async fn read_matrix_view( + &mut self, + addr: u32, + view: Option, + axis: MatrixAccessAxis, + ) -> QuantTensor { + let layout = view.map_or_else(|| self.mram.default_layout(), |view| view.layout()); + // Existing Matrix arithmetic consumes one tile. Multi-tile packet + // consumers are introduced separately; silently dropping tile_count + // here would make a plausible but incorrect result. + assert_eq!( + layout.tile_count, 1, + "a scalar Matrix opcode cannot consume a multi-tile view" + ); + let (viewed, _service) = self.mram.read_layout_tile_axis(addr, layout, 0, axis).await; + if layout.rows == self.mlen && layout.cols == self.mlen { + return viewed; + } + + assert!(layout.rows <= self.mlen && layout.cols <= self.mlen); + let padded = Tensor::zeros( + [self.mlen as i64, self.mlen as i64], + (viewed.as_tensor().kind(), tch::Device::Cpu), + ); + let source = viewed + .as_tensor() + .view([layout.rows as i64, layout.cols as i64]); + let mut destination = padded.i((0..layout.rows as i64, 0..layout.cols as i64)); + destination.copy_(&source); + QuantTensor::quantize(padded.flatten(0, -1), self.mram.ty()) + } + fn core(&self) -> MatrixCore { self.core } + #[cfg(test)] pub(crate) async fn mm(&mut self, m_addr: u32, v_addr: u32) { + self.mm_with_view(m_addr, v_addr, None).await; + } + + pub(crate) async fn mm_with_view( + &mut self, + m_addr: u32, + v_addr: u32, + view: Option, + ) { let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); // Row-granular projection ABI: the M_MM column stride is `blen * mlen` // (compiler e852c88, isa_matrix.py vram_sub_projection*), so the within-tile @@ -112,7 +156,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(mat_offset < self.mlen); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Row) + .await; // Slice columns instead of rows: [mlen, blen] let mat = full_mat .as_tensor() @@ -143,7 +189,13 @@ impl MatrixMachine { self.m_accum += vec_f32.matmul(&mat_f32); } - pub(crate) async fn bmm(&mut self, m_addr: u32, v_addr: u32, bmm_scale: f32) { + pub(crate) async fn bmm( + &mut self, + m_addr: u32, + v_addr: u32, + bmm_scale: f32, + view: Option, + ) { assert!(self.broadcast_amount * self.hlen == self.mlen); // Load matrix from matrix SRAM. let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.blen); @@ -151,7 +203,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(head_offset.is_multiple_of(self.hlen)); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Row) + .await; // Slice columns instead of rows: [hlen, mlen] let mat = full_mat @@ -202,7 +256,13 @@ impl MatrixMachine { tracing::trace!("hm_accum = {}", self.hm_accum); } - pub(crate) async fn bmv(&mut self, m_addr: u32, v_addr: u32, bmm_scale: f32) { + pub(crate) async fn bmv( + &mut self, + m_addr: u32, + v_addr: u32, + bmm_scale: f32, + view: Option, + ) { assert!(self.broadcast_amount * self.hlen == self.mlen); // Load matrix from matrix SRAM. let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.blen); @@ -210,7 +270,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(head_offset.is_multiple_of(self.hlen)); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Row) + .await; // Slice columns instead of rows: [hlen, mlen] let mat = full_mat @@ -260,7 +322,13 @@ impl MatrixMachine { tracing::trace!("hv_accum = {}", self.hv_accum); } - pub(crate) async fn btmm(&mut self, m_addr: u32, v_addr: u32, bmm_scale: f32) { + pub(crate) async fn btmm( + &mut self, + m_addr: u32, + v_addr: u32, + bmm_scale: f32, + view: Option, + ) { assert!(self.broadcast_amount * self.hlen == self.mlen); // Load matrix from matrix SRAM. let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); @@ -268,7 +336,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(head_offset.is_multiple_of(self.hlen)); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Column) + .await; // Slice columns instead of rows: [hlen, mlen] let mat = full_mat @@ -326,7 +396,13 @@ impl MatrixMachine { self.hm_accum += result_tensor; } - pub(crate) async fn btmv(&mut self, m_addr: u32, v_addr: u32, bmm_scale: f32) { + pub(crate) async fn btmv( + &mut self, + m_addr: u32, + v_addr: u32, + bmm_scale: f32, + view: Option, + ) { assert!(self.broadcast_amount * self.hlen == self.mlen); // Load matrix from matrix SRAM. let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); @@ -334,7 +410,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(head_offset.is_multiple_of(self.hlen)); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Column) + .await; // Slice columns instead of rows: [mlen, hlen] let mat = full_mat @@ -391,11 +469,18 @@ impl MatrixMachine { self.hv_accum += result_tensor; } - pub(crate) async fn tmm(&mut self, v_addr: u32, m_addr: u32) { + pub(crate) async fn tmm( + &mut self, + v_addr: u32, + m_addr: u32, + view: Option, + ) { let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); let mat_offset = assert_multiple_of(mat_offset, self.mlen); assert!(mat_offset.is_multiple_of(self.blen)); - let full_mat = self.mram.read(mat_base).await; + let full_mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Column) + .await; // Transpose then slice columns: [mlen, blen] let mat = full_mat .as_tensor() @@ -445,6 +530,49 @@ impl MatrixMachine { self.m_accum = Tensor::zeros([self.blen as i64, self.blen as i64], ACCUM_OPTS); } + /// Flush the ordinary Matrix accumulator directly into Matrix SRAM. + pub(crate) async fn mview_wo( + &mut self, + matrix_base: u32, + logical_offset: u32, + view: MatrixViewDescriptor, + ) -> sram::matrix::MatrixPacketService { + assert!( + view.shape.rows <= self.blen, + "Matrix-view writeback has {} live rows but BLEN is {}", + view.shape.rows, + self.blen + ); + assert!( + view.shape.cols >= self.blen && view.shape.cols.is_multiple_of(self.blen), + "Matrix-view output rows must contain whole BLEN-wide accumulator blocks" + ); + assert!( + logical_offset.is_multiple_of(self.blen), + "Matrix-view accumulator writeback must start at a BLEN-wide fragment boundary" + ); + self.core().compute(1).await; + let tensor = self + .m_accum + .narrow(0, 0, i64::from(view.shape.rows)) + .flatten(0, -1) + .contiguous(); + let tensor = QuantTensor::quantize(tensor, self.mram.ty()); + let service = self + .mram + .write_layout_microtile( + matrix_base, + view.layout(), + logical_offset, + tensor, + view.shape.rows, + self.blen, + ) + .await; + self.m_accum = Tensor::zeros([self.blen as i64, self.blen as i64], ACCUM_OPTS); + service + } + pub(crate) async fn bmm_wo(&mut self, v_addr: u32) { let (vec_base, vec_offset) = multiple_and_offset(v_addr, self.mlen); assert!(vec_offset.is_multiple_of(self.mlen)); @@ -490,7 +618,12 @@ impl MatrixMachine { self.hv_accum = Tensor::zeros([self.broadcast_amount as i64, self.mlen as i64], ACCUM_OPTS); } - pub(crate) async fn mv(&mut self, m_addr: u32, v_addr: u32) { + pub(crate) async fn mv( + &mut self, + m_addr: u32, + v_addr: u32, + view: Option, + ) { let (mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); tracing::debug!("======================== MV =========================="); tracing::debug!("m_addr = {:?}", m_addr); @@ -499,7 +632,9 @@ impl MatrixMachine { assert!(mat_offset.is_multiple_of(self.blen)); assert!(mat_offset < self.mlen); - let mat = self.mram.read(mat_base).await; + let mat = self + .read_matrix_view(mat_base, view, MatrixAccessAxis::Row) + .await; let vec = self.vram.read(v_addr).await; self.core().compute(self.mlen).await; // vec @ mat: [1, mlen] @ [mlen, mlen] = [1, mlen], then squeeze @@ -517,14 +652,21 @@ impl MatrixMachine { self.v_accum += result; } - pub(crate) async fn tmv(&mut self, m_addr: u32, v_addr: u32) { + pub(crate) async fn tmv( + &mut self, + m_addr: u32, + v_addr: u32, + view: Option, + ) { // TODO: `_mat_base` is computed for the assertion below but the read // uses `m_addr` directly. For tile-aligned reads they're equivalent // (integer division), but other matrix ops here use `mat_base`. Worth // investigating whether this should be `mram.read(mat_base)`. let (_mat_base, mat_offset) = multiple_and_offset(m_addr, self.mlen * self.mlen); assert!(mat_offset.is_multiple_of(self.blen)); - let mat = self.mram.read(m_addr).await; + let mat = self + .read_matrix_view(m_addr, view, MatrixAccessAxis::Column) + .await; let vec = self.vram.read(v_addr).await; self.core().compute(self.mlen).await; // vec @ transpose(mat): [1, mlen] @ [mlen, mlen] = [1, mlen], then squeeze @@ -699,4 +841,76 @@ mod tests { assert!(a0.equal(&Tensor::from_slice(&[1.0f32, 2.0, 0.0, 0.0]))); assert!(a1.equal(&Tensor::from_slice(&[5.0f32, 6.0, 0.0, 0.0]))); } + + #[tokio::test] + async fn transposed_matrix_ops_match_mathematical_transpose_with_wide_banks() { + let executor = Executor::new(); + let matrix = [ + 1.0, 2.0, 3.0, 4.0, // + 5.0, 6.0, 7.0, 8.0, // + 9.0, 10.0, 11.0, 12.0, // + 13.0, 14.0, 15.0, 16.0, + ]; + + let tmv_mram = Arc::new(MatrixSram::with_banks(4, 64, 2, bf16_plain())); + let tmv_vram = Arc::new(VectorSram::from_mx_type(4, 64, bf16_plain())); + tmv_mram.write(0, quant(&matrix)).await; + tmv_vram.write(0, quant(&[1.0, 2.0, 3.0, 4.0])).await; + let mut tmv_machine = MatrixMachine::new_with_core( + tmv_mram, + tmv_vram.clone(), + 4, + 2, + 2, + 2, + MatrixCoreProfile::big_default(), + ); + + let tmm_mram = Arc::new(MatrixSram::with_banks(4, 64, 2, bf16_plain())); + let tmm_vram = Arc::new(VectorSram::from_mx_type(4, 64, bf16_plain())); + tmm_mram.write(0, quant(&matrix)).await; + tmm_vram.write(0, quant(&[1.0, 2.0, 3.0, 4.0])).await; + tmm_vram.write(4, quant(&[2.0, 0.0, 1.0, 3.0])).await; + let mut tmm_machine = MatrixMachine::new_with_core( + tmm_mram, + tmm_vram.clone(), + 4, + 2, + 2, + 2, + MatrixCoreProfile::big_default(), + ); + + executor.spawn(async move { + tmv_machine.tmv(0, 0, None).await; + tmv_machine.mv_wo(8).await; + }); + executor.spawn(async move { + tmm_machine.tmm(0, 0, None).await; + tmm_machine.mm_wo(8, 1).await; + }); + executor.enter(Instant::ETERNITY).await; + + assert!( + tmv_vram + .read(8) + .await + .as_tensor() + .equal(&Tensor::from_slice(&[30.0f32, 70.0, 0.0, 0.0])) + ); + assert!( + tmm_vram + .read(8) + .await + .as_tensor() + .equal(&Tensor::from_slice(&[30.0f32, 70.0, 0.0, 0.0])) + ); + assert!( + tmm_vram + .read(12) + .await + .as_tensor() + .equal(&Tensor::from_slice(&[17.0f32, 41.0, 0.0, 0.0])) + ); + } } diff --git a/transactional_emulator/src/op.rs b/transactional_emulator/src/op.rs index e8c31993..7c8784fa 100644 --- a/transactional_emulator/src/op.rs +++ b/transactional_emulator/src/op.rs @@ -8,6 +8,7 @@ pub enum MatrixPrecision { pub enum VectorPrecision { Activation, KeyValue, + State, } #[derive(Debug, Clone, Copy)] @@ -16,6 +17,45 @@ pub enum VectorOrder { Reverse, } +/// Model-independent algebraic forms executed over configured Matrix views. +/// Loop bounds and broadcasting come from the views, never from a model ID. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum LTilePrimitive { + ScaleAccum, + DotReduce, + OuterUpdate, +} + +impl TryFrom for LTilePrimitive { + type Error = (); + + fn try_from(value: u8) -> Result { + match value { + 0 => Ok(Self::ScaleAccum), + 1 => Ok(Self::DotReduce), + 2 => Ok(Self::OuterUpdate), + _ => Err(()), + } + } +} + +/// Logical line direction selected independently for each `L_TILE_EXEC` input. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum LTileAxis { + Row, + Column, +} + +impl From for LTileAxis { + fn from(value: u8) -> Self { + match value { + 0 => Self::Row, + 1 => Self::Column, + _ => unreachable!("L_TILE axis is one bit"), + } + } +} + #[allow(non_camel_case_types)] #[derive(Debug)] pub enum Opcode { @@ -23,18 +63,22 @@ pub enum Opcode { M_MM { rs1: u8, rs2: u8, + view: Option, }, M_TMM { rs1: u8, rs2: u8, + view: Option, }, M_BMM { rs1: u8, rs2: u8, + view: Option, }, M_BTMM { rs1: u8, rs2: u8, + view: Option, }, M_BMM_WO { rd: u8, @@ -44,24 +88,29 @@ pub enum Opcode { rd: u8, rstride: u8, imm: u32, + view: Option, }, M_MV { rs1: u8, rs2: u8, + view: Option, }, M_TMV { rs1: u8, rs2: u8, + view: Option, }, M_BMV { rs1: u8, rs2: u8, rd: u8, + view: Option, }, M_BTMV { rs1: u8, rs2: u8, rd: u8, + view: Option, }, M_MV_WO { rd: u8, @@ -76,6 +125,7 @@ pub enum Opcode { rs1: u8, rs2: u8, rmask: u8, + view_mask: u8, }, V_ADD_VF { rd: u8, @@ -88,6 +138,7 @@ pub enum Opcode { rs1: u8, rs2: u8, rmask: u8, + view_mask: u8, }, V_SUB_VF { rd: u8, @@ -101,6 +152,7 @@ pub enum Opcode { rs1: u8, rs2: u8, rmask: u8, + view_mask: u8, }, V_MUL_VF { rd: u8, @@ -301,6 +353,35 @@ pub enum Opcode { rs2: u8, }, C_BREAK, + H_PREFETCH_V_MV { + rd: u8, + rs1: u8, + rs2: u8, + rstride: u8, + precision: VectorPrecision, + view: u8, + }, + H_STORE_V_MV { + rd: u8, + rs1: u8, + rs2: u8, + rstride: u8, + precision: VectorPrecision, + view: u8, + }, + L_TILE_CFG { + shape: u8, + mapping: u8, + slot: u8, + }, + L_TILE_EXEC { + rd: u8, + rs1: u8, + rs2: u8, + primitive: LTilePrimitive, + source_axis: LTileAxis, + scale_axis: LTileAxis, + }, } const OPERAND_WIDTH: u32 = 4; @@ -313,6 +394,49 @@ const fn mask(width: u32) -> u32 { } impl Opcode { + fn matrix_view_from(funct1: u8) -> Option { + match funct1 { + 0 => None, + 1..=4 => Some(funct1 - 1), + _ => unreachable!("caller must reject reserved Matrix-view selector"), + } + } + + fn matrix_writeback_view(immediate: u32) -> (u32, Option) { + const VIEW_MARKER: u32 = 1 << 17; + if immediate & VIEW_MARKER == 0 { + (immediate, None) + } else { + ( + immediate & ((1 << 15) - 1), + Some(((immediate >> 15) & 0x3) as u8), + ) + } + } + + fn matrix_view_vector_precision_from(funct1: u8) -> Option { + // Explicit Matrix-view DMA adds an independent state selector. Ordinary + // DMA preserves its original 0=activation/nonzero=KV interpretation. + match funct1 { + 0 => Some(VectorPrecision::Activation), + 1 => Some(VectorPrecision::KeyValue), + 2 => Some(VectorPrecision::State), + _ => None, + } + } + + fn matrix_view_dma_slot(instr: u32) -> Result, ()> { + let high = (instr >> 26) as u8; + if high == 0 { + return Ok(None); + } + // bit31 marks the form, bits30:29 are the slot, bits28:26 are zero. + if high & 0b100000 == 0 || high & 0b000111 != 0 { + return Err(()); + } + Ok(Some((high >> 3) & 0b11)) + } + #[inline] fn matrix_precision_from(funct1: u8) -> MatrixPrecision { if funct1 == 0 { @@ -353,9 +477,17 @@ impl Opcode { match opcode { 0x00 => Self::Invalid, // Matrix Operations - 0x01 => Self::M_MM { rs1, rs2 }, - 0x02 => Self::M_TMM { rs1, rs2 }, - 0x03 => { + 0x01 if funct1 <= 4 => Self::M_MM { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + }, + 0x02 if funct1 <= 4 => Self::M_TMM { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + }, + 0x03 if funct1 <= 4 => { // ISA spec defines matrix address as `gp_reg + gp_reg` but // this emulator only consumes `rs1`. M_BMV/M_BTMV honor `rd`; until // M_BMM/M_BTMM follow suit, refuse encodings that would otherwise @@ -364,34 +496,69 @@ impl Opcode { rd, 0, "M_BMM rd must be 0: emulator does not honor the spec's `gp_reg` matrix offset" ); - Self::M_BMM { rs1, rs2 } + Self::M_BMM { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + } } - 0x04 => { + 0x04 if funct1 <= 4 => { assert_eq!( rd, 0, "M_BTMM rd must be 0: emulator does not honor the spec's `gp_reg` matrix offset" ); - Self::M_BTMM { rs1, rs2 } + Self::M_BTMM { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + } } 0x05 => Self::M_BMM_WO { rd, imm: imm2 }, - 0x06 => Self::M_MM_WO { + 0x06 => { + let (imm, view) = Self::matrix_writeback_view(imm2); + Self::M_MM_WO { + rd, + rstride: rs1, + imm, + view, + } + } + 0x07 if funct1 <= 4 => Self::M_MV { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + }, + 0x08 if funct1 <= 4 => Self::M_TMV { + rs1, + rs2, + view: Self::matrix_view_from(funct1), + }, + 0x09 if funct1 <= 4 => Self::M_BMV { + rs1, + rs2, + rd, + view: Self::matrix_view_from(funct1), + }, + 0x0A if funct1 <= 4 => Self::M_BTMV { + rs1, + rs2, rd, - rstride: rs1, - imm: imm2, + view: Self::matrix_view_from(funct1), }, - 0x07 => Self::M_MV { rs1, rs2 }, - 0x08 => Self::M_TMV { rs1, rs2 }, - 0x09 => Self::M_BMV { rs1, rs2, rd }, - 0x0A => Self::M_BTMV { rs1, rs2, rd }, + 0x01..=0x04 | 0x07..=0x0A => { + tracing::error!(instr, funct1, "reserved Matrix-view selector"); + Self::Invalid + } 0x0B => Self::M_MV_WO { rd, imm: imm2 }, 0x0C => Self::M_BMV_WO { rd, imm: imm2 }, // Vector Operations - 0x0D => Self::V_ADD_VV { + 0x0D if funct1 == 0 || funct1 >= 9 => Self::V_ADD_VV { rd, rs1, rs2, rmask: rs3, + view_mask: funct1 & 7, }, 0x0E => Self::V_ADD_VF { rd, @@ -399,11 +566,12 @@ impl Opcode { rs2, rmask: rs3, }, - 0x0F => Self::V_SUB_VV { + 0x0F if funct1 == 0 || funct1 >= 9 => Self::V_SUB_VV { rd, rs1, rs2, rmask: rs3, + view_mask: funct1 & 7, }, 0x10 => Self::V_SUB_VF { rd, @@ -412,11 +580,12 @@ impl Opcode { rmask: rs3, rorder: Self::vector_order_from(funct1), }, - 0x11 => Self::V_MUL_VV { + 0x11 if funct1 == 0 || funct1 >= 9 => Self::V_MUL_VV { rd, rs1, rs2, rmask: rs3, + view_mask: funct1 & 7, }, 0x12 => Self::V_MUL_VF { rd, @@ -463,6 +632,8 @@ impl Opcode { rmask: rs3, }, + 0x0D | 0x0F | 0x11 => Self::Invalid, + // Scalar Operations (Floating-Point) 0x17 => Self::S_ADD_FP { rd, rs1, rs2 }, 0x18 => Self::S_SUB_FP { rd, rs1, rs2 }, @@ -492,20 +663,60 @@ impl Opcode { precision: Self::matrix_precision_from(funct1), }, // 0x29 => Self::H_PREFETCH_M { rd, rs1, rs2, rstride: rs3, precision: MatrixPrecision::KeyValue }, - 0x29 => Self::H_PREFETCH_V { - rd, - rs1, - rs2, - rstride: rs3, - precision: Self::vector_precision_from(funct1), + 0x29 => match Self::matrix_view_dma_slot(instr) { + Ok(None) => Self::H_PREFETCH_V { + rd, + rs1, + rs2, + rstride: rs3, + precision: Self::vector_precision_from(funct1), + }, + Ok(Some(view)) => match Self::matrix_view_vector_precision_from(funct1) { + Some(precision) => Self::H_PREFETCH_V_MV { + rd, + rs1, + rs2, + rstride: rs3, + precision, + view, + }, + None => { + tracing::error!(instr, funct1, "reserved Matrix-view prefetch precision"); + Self::Invalid + } + }, + Err(()) => { + tracing::error!(instr, funct1, "reserved H_PREFETCH_V encoding"); + Self::Invalid + } }, // 0x2A => Self::H_PREFETCH_V { rd, rs1, rs2, rstride: rs3, precision: VectorPrecision::KeyValue }, - 0x2A => Self::H_STORE_V { - rd, - rs1, - rs2, - rstride: rs3, - precision: Self::vector_precision_from(funct1), + 0x2A => match Self::matrix_view_dma_slot(instr) { + Ok(None) => Self::H_STORE_V { + rd, + rs1, + rs2, + rstride: rs3, + precision: Self::vector_precision_from(funct1), + }, + Ok(Some(view)) => match Self::matrix_view_vector_precision_from(funct1) { + Some(precision) => Self::H_STORE_V_MV { + rd, + rs1, + rs2, + rstride: rs3, + precision, + view, + }, + None => { + tracing::error!(instr, funct1, "reserved Matrix-view store precision"); + Self::Invalid + } + }, + Err(()) => { + tracing::error!(instr, funct1, "reserved H_STORE_V encoding"); + Self::Invalid + } }, // 0x2B => Self::H_STORE_V { rd, rs1, rs2, rstride: rs3, precision: VectorPrecision::KeyValue }, 0x2B => Self::C_SET_ADDR_REG { rd, rs1, rs2 }, @@ -524,6 +735,40 @@ impl Opcode { // 0x35..=0x37 (V_MAX_VF/V_MIN_VF/V_TOPK) are decoded with the other // masked vector ops above. 0x38 => Self::C_SET_TOPK_REG { rd }, + 0x3F if funct1 == 1 => { + if instr >> 26 != 0 || rs2 >= 4 || rs3 != 0 { + tracing::error!(instr, "non-canonical L_TILE_CFG encoding"); + Self::Invalid + } else { + Self::L_TILE_CFG { + shape: rd, + mapping: rs1, + slot: rs2, + } + } + } + 0x3F if funct1 == 3 => { + if instr >> 28 != 0 { + tracing::error!(instr, "non-canonical L_TILE_EXEC encoding"); + Self::Invalid + } else if let Ok(primitive) = LTilePrimitive::try_from(rs3) { + Self::L_TILE_EXEC { + rd, + rs1, + rs2, + primitive, + source_axis: LTileAxis::from(((instr >> 26) & 1) as u8), + scale_axis: LTileAxis::from(((instr >> 27) & 1) as u8), + } + } else { + tracing::error!(instr, rs3, "reserved L_TILE primitive"); + Self::Invalid + } + } + 0x3F => { + tracing::error!(instr, funct1, "reserved L_TILE form"); + Self::Invalid + } _ => { tracing::error!("Unknown opcode {opcode:#x}"); Self::Invalid @@ -551,7 +796,8 @@ mod tests { rs1, rs2, rmask, - } => assert_eq!((rd, rs1, rs2, rmask), (1, 2, 3, 4)), + view_mask, + } => assert_eq!((rd, rs1, rs2, rmask, view_mask), (1, 2, 3, 4, 0)), other => panic!("expected V_ADD_VV, got {other:?}"), } } @@ -560,7 +806,7 @@ mod tests { fn test_decode_two_register_matrix_op() { // M_MM consumes only rs1 and rs2. match Opcode::decode(rform(0x01, 0, 5, 6, 0, 0)) { - Opcode::M_MM { rs1, rs2 } => assert_eq!((rs1, rs2), (5, 6)), + Opcode::M_MM { rs1, rs2, .. } => assert_eq!((rs1, rs2), (5, 6)), other => panic!("expected M_MM, got {other:?}"), } } @@ -568,8 +814,8 @@ mod tests { #[test] fn test_decode_invalid_and_unknown_are_invalid() { assert!(matches!(Opcode::decode(0x00), Opcode::Invalid)); - // 0x3F is past the highest defined opcode. - assert!(matches!(Opcode::decode(0x3F), Opcode::Invalid)); + // V_PS_V remains declared but deliberately unimplemented. + assert!(matches!(Opcode::decode(0x31), Opcode::Invalid)); } #[test] @@ -639,7 +885,7 @@ mod tests { #[test] fn test_decode_m_bmm_rd_zero_ok() { match Opcode::decode(rform(0x03, 0, 7, 8, 0, 0)) { - Opcode::M_BMM { rs1, rs2 } => assert_eq!((rs1, rs2), (7, 8)), + Opcode::M_BMM { rs1, rs2, .. } => assert_eq!((rs1, rs2), (7, 8)), other => panic!("expected M_BMM, got {other:?}"), } } @@ -713,9 +959,14 @@ mod tests { #[test] fn test_decode_m_mm_wo_carries_rstride_and_imm2() { // M_MM_WO packs rd, rstride (= rs1 field), and the 18-bit imm2. - match Opcode::decode(i2form(0x06, 5, 6, 0x2BEEF)) { - Opcode::M_MM_WO { rd, rstride, imm } => { - assert_eq!((rd, rstride, imm), (5, 6, 0x2BEEF)) + match Opcode::decode(i2form(0x06, 5, 6, 0x0BEEF)) { + Opcode::M_MM_WO { + rd, + rstride, + imm, + view, + } => { + assert_eq!((rd, rstride, imm, view), (5, 6, 0x0BEEF, None)) } other => panic!("expected M_MM_WO, got {other:?}"), } @@ -725,7 +976,7 @@ mod tests { fn test_decode_m_bmv_carries_rd() { // M_BMV honors rd (unlike M_BMM); decode keeps all three. match Opcode::decode(rform(0x09, 9, 7, 8, 0, 0)) { - Opcode::M_BMV { rs1, rs2, rd } => assert_eq!((rs1, rs2, rd), (7, 8, 9)), + Opcode::M_BMV { rs1, rs2, rd, .. } => assert_eq!((rs1, rs2, rd), (7, 8, 9)), other => panic!("expected M_BMV, got {other:?}"), } } @@ -748,6 +999,15 @@ mod tests { .. } )); + for legacy_nonzero in 2..=15 { + assert!(matches!( + Opcode::decode(rform(0x29, 0, 0, 0, 0, legacy_nonzero)), + Opcode::H_PREFETCH_V { + precision: VectorPrecision::KeyValue, + .. + } + )); + } } #[test] @@ -762,6 +1022,48 @@ mod tests { } => assert_eq!((rd, rs1, rs2, rstride), (1, 2, 3, 4)), other => panic!("expected H_STORE_V KeyValue, got {other:?}"), } + for legacy_nonzero in 2..=15 { + assert!(matches!( + Opcode::decode(rform(0x2A, 1, 2, 3, 4, legacy_nonzero)), + Opcode::H_STORE_V { + precision: VectorPrecision::KeyValue, + .. + } + )); + } + } + + #[test] + fn test_decode_vector_dma_matrix_view_form_and_reject_reserved_high_bits() { + let marker = 1_u32 << 31; + let slot = 3_u32 << 29; + match Opcode::decode(rform(0x29, 1, 2, 3, 4, 2) | marker | slot) { + Opcode::H_PREFETCH_V_MV { + rd, + rs1, + rs2, + rstride, + precision: VectorPrecision::State, + view, + } => assert_eq!((rd, rs1, rs2, rstride, view), (1, 2, 3, 4, 3)), + other => panic!("expected Matrix-view prefetch, got {other:?}"), + } + assert!(matches!( + Opcode::decode(rform(0x2A, 1, 2, 3, 4, 2) | marker | (2 << 29)), + Opcode::H_STORE_V_MV { + precision: VectorPrecision::State, + view: 2, + .. + } + )); + assert!(matches!( + Opcode::decode(rform(0x29, 1, 2, 3, 4, 2) | marker | (1 << 28)), + Opcode::Invalid + )); + assert!(matches!( + Opcode::decode(rform(0x29, 1, 2, 3, 4, 3) | marker), + Opcode::Invalid + )); } // ---------- control ops ---------- @@ -787,6 +1089,139 @@ mod tests { assert!(matches!(Opcode::decode(0x34), Opcode::C_BREAK)); } + // ---------- Mamba / selective-SSM extensions ---------- + + /// Every opcode PLENA_Compiler declares must decode to something here. + /// + /// `decode` carries three separate comments saying encodings "must stay in + /// sync with PLENA_Compiler's doc/operation.svh". That was enforced by + /// author diligence alone -- exactly the arrangement that let the compiler's + /// FPRAM depth (1024) drift from the SystemVerilog's (512) unnoticed. The + /// submodule is checked out recursively in every CI job, so the header can + /// simply be read. + #[test] + fn every_compiler_opcode_decodes_to_something() { + let svh = include_str!("../../PLENA_Compiler/doc/operation.svh"); + + // Declared in the header but deliberately not modelled. A new entry is + // a decision to argue for in review, not a default. + const NOT_MODELLED: &[&str] = &[ + // The sentinel. Invalid is precisely what it must decode to. + "INVALID_OPCODE", + // Declared by PLENA but never implemented, in RTL or here. The + "V_PS_V", + // Likewise declared and unimplemented; nothing emits it. + "C_HADAMARD_TRANSFORM", + // Owned by the Shared Expert branch. Their numeric reservation is + // part of this branch's conflict-free ABI, but route execution is + // merged independently from L-Compute. + "C_ROUTE_BEGIN", + "C_ROUTE_LOOP_START", + "C_ROUTE_LOOP_END", + "V_ROUTE_MUL", + ]; + + let mut checked = 0; + for line in svh.lines() { + let line = line.trim(); + if line.starts_with("//") { + continue; + } + let Some((name, rest)) = line.split_once('=') else { + continue; + }; + let name = name.trim(); + let Some(hex) = rest.trim().strip_prefix("6'h") else { + continue; + }; + let hex: String = hex.chars().take_while(|c| c.is_ascii_hexdigit()).collect(); + let Ok(opcode) = u32::from_str_radix(&hex, 16) else { + continue; + }; + if NOT_MODELLED.contains(&name) { + continue; + } + checked += 1; + assert!( + !matches!( + Opcode::decode(rform( + opcode, + 0, + 0, + 0, + 0, + if opcode == 0x3F { 1 } else { 0 } + )), + Opcode::Invalid + ), + "{name} = 6'h{opcode:02X} is declared in PLENA_Compiler's \ + doc/operation.svh but decodes to Invalid here" + ); + } + assert!( + checked > 40, + "only {checked} opcodes parsed out of the header -- the parse broke, \ + so this guard was passing vacuously" + ); + } + + #[test] + fn l_tile_forms_and_explicit_matrix_consumer_match_compiler_words() { + match Opcode::decode(rform(0x3F, 7, 9, 2, 0, 1)) { + Opcode::L_TILE_CFG { + shape, + mapping, + slot, + } => assert_eq!((shape, mapping, slot), (7, 9, 2)), + other => panic!("expected L_TILE_CFG, got {other:?}"), + } + assert!(matches!( + Opcode::decode(rform(0x3F, 9, 2, 2, 0, 2)), + Opcode::Invalid + )); + match Opcode::decode(rform(0x09, 9, 5, 6, 0, 3)) { + Opcode::M_BMV { rd, rs1, rs2, view } => { + assert_eq!((rd, rs1, rs2, view), (9, 5, 6, Some(2))); + } + other => panic!("expected viewed M_BMV, got {other:?}"), + } + match Opcode::decode(i2form(0x06, 4, 0, (1 << 17) | (2 << 15) | 5)) { + Opcode::M_MM_WO { + rd, + rstride, + imm, + view, + } => assert_eq!((rd, rstride, imm, view), (4, 0, 5, Some(2))), + other => panic!("expected viewed M_MM_WO, got {other:?}"), + } + match Opcode::decode(rform(0x0D, 4, 5, 6, 0, 0x8 | 0b110)) { + Opcode::V_ADD_VV { + rd, + rs1, + rs2, + view_mask, + .. + } => assert_eq!((rd, rs1, rs2, view_mask), (4, 5, 6, 0b110)), + other => panic!("expected Matrix-view V_ADD_VV, got {other:?}"), + } + } + + #[test] + fn l_mview_rejects_reserved_bits_forms_and_consumer_slots() { + assert!(matches!( + Opcode::decode(rform(0x3F, 7, 9, 2, 1, 1)), + Opcode::Invalid + )); + assert!(matches!( + Opcode::decode(rform(0x3F, 7, 9, 2, 0, 6)), + Opcode::Invalid + )); + assert!(matches!( + Opcode::decode(rform(0x01, 0, 5, 6, 0, 5)), + Opcode::Invalid + )); + } + #[test] fn test_decode_v_shft_v() { match Opcode::decode(rform(0x32, 1, 2, 3, 0, 0)) { @@ -800,7 +1235,7 @@ mod tests { // M_BTMV (unlike M_BTMM) honors rd; decode keeps all three fields, and // unlike M_BMM/M_BTMM it does not assert rd == 0. match Opcode::decode(rform(0x0A, 9, 7, 8, 0, 0)) { - Opcode::M_BTMV { rs1, rs2, rd } => assert_eq!((rs1, rs2, rd), (7, 8, 9)), + Opcode::M_BTMV { rs1, rs2, rd, .. } => assert_eq!((rs1, rs2, rd), (7, 8, 9)), other => panic!("expected M_BTMV, got {other:?}"), } } @@ -826,20 +1261,12 @@ mod tests { rs1, rs2, rmask, + .. } => assert_eq!((rd, rs1, rs2, rmask), (15, 15, 15, 15)), other => panic!("expected V_ADD_VV, got {other:?}"), } } - #[test] - fn test_decode_funct1_does_not_bleed_into_rmask() { - // funct1 (bits 22..26) must not leak into rmask (= rs3, bits 18..22). - match Opcode::decode(rform(0x0D, 0, 0, 0, 0, 0xF)) { - Opcode::V_ADD_VV { rmask, .. } => assert_eq!(rmask, 0), - other => panic!("expected V_ADD_VV, got {other:?}"), - } - } - #[test] fn test_decode_vector_scalar_minmax() { match Opcode::decode(rform(0x35, 1, 2, 3, 4, 0)) { @@ -848,6 +1275,7 @@ mod tests { rs1, rs2, rmask, + .. } => { assert_eq!((rd, rs1, rs2, rmask), (1, 2, 3, 4)); } @@ -859,6 +1287,7 @@ mod tests { rs1, rs2, rmask, + .. } => { assert_eq!((rd, rs1, rs2, rmask), (5, 6, 7, 8)); } diff --git a/transactional_emulator/src/runner.rs b/transactional_emulator/src/runner.rs index 04a129c1..ae0d772e 100644 --- a/transactional_emulator/src/runner.rs +++ b/transactional_emulator/src/runner.rs @@ -12,7 +12,7 @@ use crate::matrix_core::MatrixCoreProfile; use crate::matrix_machine::MatrixMachine; use crate::runtime_config::{ BLEN, BROADCAST_AMOUNT, HBM_SIZE, HLEN, MATRIX_SRAM_SIZE, MATRIX_SRAM_TYPE, - MAX_LOOP_INSTRUCTIONS, MLEN, PREFETCH_M_AMOUNT, PREFETCH_V_AMOUNT, STORE_V_AMOUNT, + MAX_LOOP_INSTRUCTIONS, MLEN, PREFETCH_M_AMOUNT, PREFETCH_V_AMOUNT, STATE_TYPE, STORE_V_AMOUNT, VECTOR_SRAM_SIZE, VECTOR_SRAM_TYPE, VLEN, }; use crate::stage_profile::StageProfiler; @@ -96,6 +96,7 @@ pub(crate) async fn run_from_cli() { vector_sram_size = *VECTOR_SRAM_SIZE, matrix_type = ?*MATRIX_SRAM_TYPE, vector_type = ?*VECTOR_SRAM_TYPE, + recurrent_state_type = ?*STATE_TYPE, "SRAM" ); tracing::info!( @@ -111,7 +112,12 @@ pub(crate) async fn run_from_cli() { "Config source" ); - let mram = Arc::new(MatrixSram::new(*MLEN, *MATRIX_SRAM_SIZE, *MATRIX_SRAM_TYPE)); // Matrix SRAM + let mram = Arc::new(MatrixSram::with_banks( + *MLEN, + *MATRIX_SRAM_SIZE, + *BLEN, + *MATRIX_SRAM_TYPE, + )); // Matrix SRAM let vram = Arc::new(VectorSram::from_mx_type( *VLEN, *VECTOR_SRAM_SIZE, @@ -239,6 +245,17 @@ pub(crate) async fn run_from_cli() { .do_ops(&decoded_ops, stage_profiler.as_mut(), timing_driver) .await; + let matrix_packet = accelerator.matrix_view_packet_counters(); + tracing::info!( + packets = matrix_packet.packets, + values = matrix_packet.values, + bank_words = matrix_packet.bank_words, + service_cycles = matrix_packet.service_cycles, + ideal_cycles = matrix_packet.ideal_cycles, + bank_stall_cycles = matrix_packet.bank_stall_cycles, + "Matrix-view packet counters" + ); + let serial_duration = Executor::current().now() - Instant::INIT; if let Some(sb) = scoreboard.as_ref() { let stats = sb.stats; @@ -297,17 +314,22 @@ pub(crate) async fn run_from_cli() { let intsram_bytes = accelerator.intsram_dump_bytes(); dump_to_file("intsram_dump.bin", &intsram_bytes); - // Dump HBM — skipped unless DEBUG tracing is enabled because HBM_SIZE may - // be 128 GiB+. Tests run with --log-level warn and don't need hbm_dump.bin; - // only manual debug runs dump HBM. - if tracing::enabled!(tracing::Level::DEBUG) { + // Dump HBM only on an explicit request or under DEBUG tracing because the + // modeled capacity may be 128 GiB+. The explicit path lets connected + // numerical tests inspect state/output writes without enabling noisy logs. + if opts.hbm_dump.is_some() || tracing::enabled!(tracing::Level::DEBUG) { let hbm_size = effective_hbm_size; let mut hbm_bytes = vec![0u8; hbm_size]; hbm.model().data().with_data(|f| { let len = std::cmp::min(hbm_size, f.len()); hbm_bytes[..len].copy_from_slice(&f[..len]); }); - dump_to_file("hbm_dump.bin", &hbm_bytes); + let path = opts + .hbm_dump + .as_deref() + .and_then(|path| path.to_str()) + .unwrap_or("hbm_dump.bin"); + dump_to_file(path, &hbm_bytes); } let memory_stats = hbm.statistics(); diff --git a/transactional_emulator/src/runtime_config.rs b/transactional_emulator/src/runtime_config.rs index ff3450ef..8db4fd83 100644 --- a/transactional_emulator/src/runtime_config.rs +++ b/transactional_emulator/src/runtime_config.rs @@ -40,6 +40,7 @@ pub(crate) static MATRIX_KV_TYPE: LazyLock = LazyLock::new(matrix_kv pub(crate) static VECTOR_ACTIVATION_TYPE: LazyLock = LazyLock::new(vector_activation_type); pub(crate) static VECTOR_KV_TYPE: LazyLock = LazyLock::new(vector_kv_type); +pub(crate) static STATE_TYPE: LazyLock = LazyLock::new(state_type); pub(crate) static PREFETCH_M_AMOUNT: LazyLock = LazyLock::new(|| { let raw = hbm_m_prefetch_amount(); let mlen = mlen(); diff --git a/transactional_emulator/src/timing.rs b/transactional_emulator/src/timing.rs index df3e236f..e9996100 100644 --- a/transactional_emulator/src/timing.rs +++ b/transactional_emulator/src/timing.rs @@ -80,6 +80,16 @@ pub(crate) fn take_charged() -> u64 { CHARGED_CYCLES.with(|c| c.replace(0)) } +/// Charge real Matrix SRAM packet service through the existing timing mode. +pub(crate) async fn charge_bank_cycles(cycles: u64) { + charge_cycles(cycles.try_into().expect("bank service exceeds u32 cycles")).await; +} + +/// Charge a packet primitive using the configured Vector arithmetic latency. +pub(crate) async fn charge_arithmetic_cycles(cycles: u32) { + charge_cycles(cycles).await; +} + #[cfg(test)] mod tests { use runtime::{Duration, Executor, Instant}; diff --git a/transactional_emulator/src/vector_machine.rs b/transactional_emulator/src/vector_machine.rs index 13a63c66..9ab5e4d2 100644 --- a/transactional_emulator/src/vector_machine.rs +++ b/transactional_emulator/src/vector_machine.rs @@ -11,7 +11,7 @@ use std::sync::Arc; use half::bf16; -use quantize::{QuantTensor, tensor_from_f32_slice}; +use quantize::{QuantTensor, tensor_from_f32_slice, tensor_to_f32_vec}; use sram::VectorSram; use tch::Tensor; @@ -29,6 +29,45 @@ pub(crate) struct VectorMachine { mask_unit: u32, } +#[derive(Clone, Copy, Debug)] +pub(crate) enum VectorBinaryOp { + Add, + Sub, + Mul, +} + +/// Coefficient representation selected by the Matrix-view descriptor. +/// A one-tile packet has the same length in either representation, so packet +/// length cannot identify whether coefficients are global or packet-local. +#[derive(Clone, Copy, Debug)] +pub(crate) enum TileScaleLayout { + Compact { first_tile: u32 }, + Expanded, +} + +impl TileScaleLayout { + fn validate(self, values: usize, tiles: usize, width: u32, per_tile: usize) { + assert!(width as usize >= per_tile); + match self { + Self::Compact { first_tile } => { + assert_eq!(values, width as usize); + assert!( + per_tile * (first_tile as usize + tiles) <= values, + "compact L_TILE coefficients must cover every addressed tile" + ); + } + Self::Expanded => assert_eq!(values, tiles * width as usize), + } + } + + fn coefficient_index(self, tile: usize, width: u32, per_tile: usize) -> usize { + match self { + Self::Compact { first_tile } => per_tile * (first_tile as usize + tile), + Self::Expanded => tile * width as usize, + } + } +} + impl VectorMachine { pub(crate) fn new(vram: Arc, tile_size: u32, mask_unit: u32) -> Self { Self { @@ -458,6 +497,148 @@ impl VectorMachine { cycle!((*VECTOR_MAX_CYCLES).saturating_mul(expert_count as u32)); (indices, weights) } + pub(crate) fn tile_size(&self) -> u32 { + self.tile_size + } + + pub(crate) async fn binary_packet( + &self, + op: VectorBinaryOp, + a: QuantTensor, + b: QuantTensor, + rmask: u8, + mask: u32, + ) -> QuantTensor { + let apply = |lhs: &Tensor, rhs: &Tensor| match op { + VectorBinaryOp::Add => lhs + rhs, + VectorBinaryOp::Sub => lhs - rhs, + VectorBinaryOp::Mul => lhs * rhs, + }; + let result = if rmask == 0 { + apply(a.as_tensor(), b.as_tensor()) + } else { + let result = a.as_tensor().shallow_clone(); + let total_heads = self.tile_size / self.mask_unit; + for head in 0..total_heads { + if (mask & (1 << head)) != 0 { + let start = (head * self.mask_unit) as i64; + let end = ((head + 1) * self.mask_unit) as i64; + let lhs = result.narrow(0, start, end - start); + let rhs = b.as_tensor().narrow(0, start, end - start); + let updated = apply(&lhs, &rhs); + result.narrow(0, start, end - start).copy_(&updated); + } + } + result + }; + match op { + VectorBinaryOp::Add | VectorBinaryOp::Sub => { + crate::timing::charge_arithmetic_cycles(*VECTOR_ADD_CYCLES).await; + } + VectorBinaryOp::Mul => { + crate::timing::charge_arithmetic_cycles(*VECTOR_MUL_CYCLES).await; + } + } + QuantTensor::quantize(result, a.data_type()) + } + + pub(crate) async fn tile_scale_accum( + &self, + destination: QuantTensor, + source: QuantTensor, + scales: QuantTensor, + row_width: u32, + scale_width: u32, + scale_layout: TileScaleLayout, + ) -> QuantTensor { + let dst = tensor_to_f32_vec(destination.as_tensor()); + let src = tensor_to_f32_vec(source.as_tensor()); + let coeff = tensor_to_f32_vec(scales.as_tensor()); + let rows = dst.len() / row_width as usize; + assert_eq!(dst.len(), rows * row_width as usize); + assert!(src.len() == row_width as usize || src.len() == dst.len()); + scale_layout.validate(coeff.len(), rows, scale_width, 2); + + let mut result = vec![0_f32; dst.len()]; + for row in 0..rows { + let coefficient = scale_layout.coefficient_index(row, scale_width, 2); + let a = coeff[coefficient]; + let b = coeff[coefficient + 1]; + for col in 0..row_width as usize { + let index = row * row_width as usize + col; + let source_index = if src.len() == row_width as usize { + col + } else { + index + }; + result[index] = a * dst[index] + b * src[source_index]; + } + } + crate::timing::charge_arithmetic_cycles(2 * *VECTOR_MUL_CYCLES + *VECTOR_ADD_CYCLES).await; + QuantTensor::quantize(tensor_from_f32_slice(&result), destination.data_type()) + } + + pub(crate) async fn tile_dot_accumulate( + &self, + accumulator: &mut [f32], + rows: QuantTensor, + scales: QuantTensor, + row_width: u32, + scale_width: u32, + scale_layout: TileScaleLayout, + ) { + let values = tensor_to_f32_vec(rows.as_tensor()); + let coeff = tensor_to_f32_vec(scales.as_tensor()); + assert!(row_width > 0); + assert!(values.len().is_multiple_of(row_width as usize)); + let row_count = values.len() / row_width as usize; + assert_eq!(accumulator.len(), values.len()); + scale_layout.validate(coeff.len(), row_count, scale_width, 1); + + for row in 0..row_count { + let scale = coeff[scale_layout.coefficient_index(row, scale_width, 1)]; + for lane in 0..row_width as usize { + let index = row * row_width as usize + lane; + accumulator[index] += values[index] * scale; + } + } + crate::timing::charge_arithmetic_cycles(*VECTOR_MUL_CYCLES + *VECTOR_ADD_CYCLES).await; + } + + pub(crate) async fn tile_outer_update( + &self, + destination: QuantTensor, + vector: QuantTensor, + scales: QuantTensor, + row_width: u32, + scale_width: u32, + scale_layout: TileScaleLayout, + ) -> QuantTensor { + let dst = tensor_to_f32_vec(destination.as_tensor()); + let source = tensor_to_f32_vec(vector.as_tensor()); + let coeff = tensor_to_f32_vec(scales.as_tensor()); + let rows = dst.len() / row_width as usize; + assert!( + source.len() == row_width as usize || source.len() == dst.len(), + "OUTER_UPDATE source must be shared by every tile or provide one row per tile" + ); + scale_layout.validate(coeff.len(), rows, scale_width, 1); + let mut result = dst; + for row in 0..rows { + let scale = coeff[scale_layout.coefficient_index(row, scale_width, 1)]; + for col in 0..row_width as usize { + let index = row * row_width as usize + col; + let source_index = if source.len() == row_width as usize { + col + } else { + index + }; + result[index] += source[source_index] * scale; + } + } + crate::timing::charge_arithmetic_cycles(*VECTOR_MUL_CYCLES + *VECTOR_ADD_CYCLES).await; + QuantTensor::quantize(tensor_from_f32_slice(&result), destination.data_type()) + } } #[cfg(test)] diff --git a/transactional_emulator/testbench/README.md b/transactional_emulator/testbench/README.md index ff6f3d6d..38c71493 100644 --- a/transactional_emulator/testbench/README.md +++ b/transactional_emulator/testbench/README.md @@ -47,3 +47,18 @@ testbench/ MoE and decoder-block bring-up lives under `routed_moe/` and `models/` so that reviewers can distinguish reusable operator coverage from model semantics harnesses. + +## Matrix SRAM views and L_TILE + +```bash +just test-matrix-lcompute +``` + +`aten/matrix_lcompute_test.py` checks projection writeback into a Matrix SRAM +view, followed by prepared Mamba/KDA recurrence through the Compiler, assembler +and Rust emulator. It covers four tokens at the model state dimensions with +fixed and phased layouts, BF16 state readback, and bank-service checks. +`aten/_matrix_lcompute.py` holds the reference formulas and input packing. + +These deterministic operator tests validate the view/L_TILE mechanism; they +do not execute complete model checkpoints or produce end-to-end speedups. diff --git a/transactional_emulator/testbench/aten/_matrix_lcompute.py b/transactional_emulator/testbench/aten/_matrix_lcompute.py new file mode 100644 index 00000000..b9a9e0e7 --- /dev/null +++ b/transactional_emulator/testbench/aten/_matrix_lcompute.py @@ -0,0 +1,330 @@ +"""Independent BF16 references and memory fixtures for Matrix L-Compute tests.""" + +from __future__ import annotations + +from collections.abc import Iterator +from contextlib import contextmanager +import os +from pathlib import Path + +import numpy as np +import tomlkit +import torch + +from compiler.aten.plena.program_lcompute import ( + BF16_BYTES, + KIMI_KDA, + NEMOTRON_MAMBA, + MatrixRecurrenceSpec, + MatrixSramPoint, + RecurrenceFieldPacket, + RecurrenceKind, + RecurrenceWorkingSet, +) + +REPO_ROOT = Path(__file__).resolve().parents[3] +SEED = 20260903 + + +def _bf16(value: torch.Tensor) -> torch.Tensor: + return value.detach().cpu().to(torch.bfloat16).float().contiguous() + + +def _bf16_bytes(value: torch.Tensor) -> bytes: + bits = _bf16(value).to(torch.bfloat16).view(torch.uint16).numpy() + return bits.astype(" torch.Tensor: + bits = np.frombuffer(image, dtype=" int: + return ((value + multiple - 1) // multiple) * multiple + + +def _write_packet( + image: bytearray, + packet: RecurrenceFieldPacket, + values: torch.Tensor, +) -> None: + flat = _bf16(values).flatten() + if flat.numel() != packet.logical_values: + raise AssertionError(f"{packet.key}: generated {flat.numel()} values, expected {packet.logical_values}") + padded = torch.zeros(packet.transfer_values, dtype=torch.float32) + padded[: flat.numel()] = flat + payload = _bf16_bytes(padded) + begin = packet.hbm_byte_offset + end = begin + len(payload) + image[begin:end] = payload + + +def _mamba_inputs(token: int, seed: int = SEED) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed + 101 * token) + heads, rows, width = ( + NEMOTRON_MAMBA.heads, + NEMOTRON_MAMBA.recurrence_rows, + NEMOTRON_MAMBA.row_elements, + ) + return { + "x": _bf16(torch.randn(heads, width, generator=generator) * 0.08), + "dt": _bf16(0.08 + torch.rand(heads, generator=generator) * 0.08), + "a": _bf16(0.82 + torch.rand(heads, rows, generator=generator) * 0.12), + "b": _bf16(torch.randn(heads, rows, generator=generator) * 0.025), + "c": _bf16(torch.randn(heads, rows, generator=generator) * 0.03), + "d": _bf16(0.15 + torch.rand(heads, generator=generator) * 0.15), + } + + +def _mamba_reference( + state: torch.Tensor, + operands: dict[str, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + x, dt = operands["x"], operands["dt"] + scratch = _bf16(dt[:, None] * x) + state = _bf16(operands["a"][:, :, None] * state + operands["b"][:, :, None] * scratch[:, None, :]) + accumulator = torch.zeros_like(x) + for row in range(state.shape[1]): + accumulator += state[:, row, :] * operands["c"][:, row, None] + output = _bf16(accumulator) + output = _bf16(output + operands["d"][:, None] * x) + return output, state + + +def _mamba_packet_values( + packet: RecurrenceFieldPacket, + operands: dict[str, torch.Tensor], + working_set: RecurrenceWorkingSet, +) -> torch.Tensor: + group_heads = working_set.group_heads + first = packet.group * group_heads + last = first + group_heads + chunk = 0 if packet.chunk is None else packet.chunk + if packet.field in {"x", "value"}: + return operands["x"][first:last] + if packet.field in {"scratch_zero", "output_zero", "output_result"}: + return torch.zeros(packet.logical_values) + if packet.field == "dt": + values = torch.zeros(packet.logical_values) + values[1 : 2 * group_heads : 2] = operands["dt"][first:last] + return values + if packet.field == "d": + values = torch.zeros(packet.logical_values) + values[0 : 2 * group_heads : 2] = 1.0 + values[1 : 2 * group_heads : 2] = operands["d"][first:last] + return values + row_first = chunk * working_set.state_rows_per_chunk + row_last = row_first + working_set.state_rows_per_chunk + if packet.field == "update": + values = torch.zeros( + working_set.state_rows_per_chunk, + working_set.allocation(packet.target).descriptor.shape.cols, + ) + values[:, 0 : 2 * group_heads : 2] = operands["a"][first:last, row_first:row_last].T + values[:, 1 : 2 * group_heads : 2] = operands["b"][first:last, row_first:row_last].T + return values + if packet.field == "c": + values = torch.zeros( + working_set.state_rows_per_chunk, + working_set.allocation(packet.target).descriptor.shape.cols, + ) + values[:, :group_heads] = operands["c"][first:last, row_first:row_last].T + return values + raise KeyError(packet.field) + + +def _kda_inputs(token: int, seed: int = SEED) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed + 1009 + 101 * token) + heads, keys, values = KIMI_KDA.heads, KIMI_KDA.recurrence_rows, KIMI_KDA.row_elements + return { + "decay": _bf16(0.84 + torch.rand(heads, keys, generator=generator) * 0.12), + "key": _bf16(torch.randn(heads, keys, generator=generator) * 0.025), + "query": _bf16(torch.randn(heads, keys, generator=generator) * 0.025), + "value": _bf16(torch.randn(heads, values, generator=generator) * 0.08), + "beta": _bf16(0.2 + torch.rand(heads, generator=generator) * 0.35), + } + + +def _kda_reference( + state: torch.Tensor, + operands: dict[str, torch.Tensor], +) -> tuple[torch.Tensor, torch.Tensor]: + state = _bf16(operands["decay"][:, :, None] * state) + prediction = torch.zeros_like(operands["value"]) + for key in range(state.shape[1]): + prediction += state[:, key, :] * operands["key"][:, key, None] + prediction = _bf16(prediction) + error = _bf16(operands["beta"][:, None] * (operands["value"] - prediction)) + state = _bf16(state + operands["key"][:, :, None] * error[:, None, :]) + output = torch.zeros_like(error) + for key in range(state.shape[1]): + output += state[:, key, :] * operands["query"][:, key, None] + return _bf16(output), state + + +def _kda_packet_values( + packet: RecurrenceFieldPacket, + operands: dict[str, torch.Tensor], + working_set: RecurrenceWorkingSet, +) -> torch.Tensor: + group_heads = working_set.group_heads + first = packet.group * group_heads + last = first + group_heads + if packet.field in {"prediction_zero", "output_zero", "output_result"}: + return torch.zeros(packet.logical_values) + if packet.field == "value": + return operands["value"][first:last] + if packet.field == "beta": + values = torch.empty(2 * group_heads) + values[0::2] = operands["beta"][first:last] + values[1::2] = -operands["beta"][first:last] + return values + chunk = 0 if packet.chunk is None else packet.chunk + row_first = chunk * working_set.state_rows_per_chunk + row_last = row_first + working_set.state_rows_per_chunk + descriptor = working_set.allocation(packet.target).descriptor + if packet.field == "decay": + values = torch.zeros(group_heads, 2, descriptor.shape.cols) + values[:, 0, : working_set.state_rows_per_chunk] = operands["decay"][first:last, row_first:row_last] + return values + if packet.field in {"key", "query"}: + values = torch.zeros(group_heads, 1, descriptor.shape.cols) + values[:, 0, : working_set.state_rows_per_chunk] = operands[packet.field][first:last, row_first:row_last] + return values + raise KeyError(packet.field) + + +def _state_seed(spec: MatrixRecurrenceSpec, seed: int = SEED) -> torch.Tensor: + generator = torch.Generator().manual_seed(seed + (0 if spec.kind is RecurrenceKind.MAMBA else 5003)) + return _bf16( + torch.randn( + spec.heads, + spec.recurrence_rows, + spec.row_elements, + generator=generator, + ) + * 0.04 + ) + + +def _pack_state_hbm(state: torch.Tensor, working_set: RecurrenceWorkingSet) -> bytes: + """Pack logical [head,row,lane] state into the Compiler's DMA packet ABI.""" + + packets = [] + for group in range(working_set.groups): + head_first = group * working_set.group_heads + head_last = head_first + working_set.group_heads + for chunk in range(working_set.chunks): + row_first = chunk * working_set.state_rows_per_chunk + row_last = row_first + working_set.state_rows_per_chunk + packets.append(state[head_first:head_last, row_first:row_last, :]) + return b"".join(_bf16_bytes(packet) for packet in packets) + + +def _unpack_state_hbm( + image: bytes, + working_set: RecurrenceWorkingSet, +) -> torch.Tensor: + """Restore logical [head,row,lane] state from the packet-major DMA ABI.""" + + spec = working_set.spec + state = torch.empty( + spec.heads, + spec.recurrence_rows, + spec.row_elements, + dtype=torch.float32, + ) + packet_values = working_set.group_heads * working_set.state_rows_per_chunk * spec.row_elements + packet_bytes = packet_values * BF16_BYTES + packet_index = 0 + for group in range(working_set.groups): + head_first = group * working_set.group_heads + head_last = head_first + working_set.group_heads + for chunk in range(working_set.chunks): + row_first = chunk * working_set.state_rows_per_chunk + row_last = row_first + working_set.state_rows_per_chunk + values = _read_bf16( + image, + packet_index * packet_bytes, + packet_values, + ).reshape( + working_set.group_heads, + working_set.state_rows_per_chunk, + spec.row_elements, + ) + state[head_first:head_last, row_first:row_last, :] = values + packet_index += 1 + return state + + +def _write_settings(build_dir: Path, point: MatrixSramPoint) -> Path: + with (REPO_ROOT / "plena_settings.toml").open() as file: + config = tomlkit.load(file) + txn = config["TRANSACTIONAL"]["CONFIG"] + txn["MLEN"]["value"] = point.mlen + txn["VLEN"]["value"] = point.mlen + txn["BLEN"]["value"] = point.bank_width + txn["HLEN"]["value"] = 128 + txn["BROADCAST_AMOUNT"]["value"] = point.bank_width + # The transactional Matrix SRAM setting is a count of MLEN-wide physical + # rows, not a count of scalar elements. At the paper point this is + # 1 MiB / (2048 values * 2 B) = 256 rows. + txn["MATRIX_SRAM_SIZE"]["value"] = point.depth_rows + # This connected program never uses ordinary Vector-SRAM operands. Keep a + # small legal instance so the mandatory post-run dump stays bounded. + txn["VECTOR_SRAM_SIZE"]["value"] = 64 + txn["HBM_V_Prefetch_Amount"]["value"] = 1 + txn["HBM_V_Writeback_Amount"]["value"] = 1 + path = build_dir / "plena_settings.toml" + with path.open("w") as file: + tomlkit.dump(config, file) + return path + + +@contextmanager +def _setting_override(path: Path) -> Iterator[None]: + previous = os.environ.get("PLENA_SETTINGS_TOML") + os.environ["PLENA_SETTINGS_TOML"] = str(path) + try: + yield + finally: + if previous is None: + os.environ.pop("PLENA_SETTINGS_TOML", None) + else: + os.environ["PLENA_SETTINGS_TOML"] = previous + + +# BF16 recurrence acceptance policy: keep the existing per-element outlier +# bound and additionally cap aggregate relative L2 error at 1%. Fixed/chunked +# reductions round at different boundaries; this is an explicit error budget, +# not a claim that their observed error is the ISA's mathematical tolerance. +# For effectively zero tensors, permit at most 1e-7 RMS absolute error so the +# norm test does not divide by zero or reject harmless sub-signal roundoff. +RECURRENCE_RELATIVE_L2_LIMIT = 1e-2 +RECURRENCE_ZERO_RMS_LIMIT = 1e-7 + + +def _assert_close(name: str, actual: torch.Tensor, expected: torch.Tensor, *, exact: bool = False) -> dict[str, float]: + if actual.shape != expected.shape: + raise AssertionError(f"{name}: shape {tuple(actual.shape)} != {tuple(expected.shape)}") + if not torch.isfinite(actual).all() or not torch.isfinite(expected).all(): + raise AssertionError(f"{name}: non-finite actual or reference values") + if exact and not torch.equal(actual.view(torch.int32), expected.view(torch.int32)): + raise AssertionError(f"{name}: exact BF16 comparison failed") + error = (actual - expected).abs() + max_abs = float(error.max()) if error.numel() else 0.0 + error_norm = torch.linalg.vector_norm(error) + expected_norm = torch.linalg.vector_norm(expected) + relative_l2 = float(error_norm / expected_norm.clamp_min(1e-12)) + norm_budget = max( + RECURRENCE_RELATIVE_L2_LIMIT * float(expected_norm), + RECURRENCE_ZERO_RMS_LIMIT * error.numel() ** 0.5, + ) + mismatch = int((error > (1e-2 + 1e-2 * expected.abs())).sum()) + if mismatch or float(error_norm) > norm_budget: + raise AssertionError( + f"{name}: {mismatch}/{actual.numel()} values mismatch; max_abs={max_abs}, " + f"relative_l2={relative_l2}, error_l2={float(error_norm)}, l2_budget={norm_budget}" + ) + return {"max_abs": max_abs, "relative_l2": relative_l2} diff --git a/transactional_emulator/testbench/aten/matrix_lcompute_test.py b/transactional_emulator/testbench/aten/matrix_lcompute_test.py new file mode 100644 index 00000000..1bd63d5f --- /dev/null +++ b/transactional_emulator/testbench/aten/matrix_lcompute_test.py @@ -0,0 +1,425 @@ +"""Compiler-to-Rust Matrix-SRAM projection and four-token recurrence checks. + +The projection must write the consumer view directly. Prepared Mamba/KDA +programs use deterministic BF16 inputs at the official state sizes, execute +four tokens in Rust, and read every output plus the final state from HBM. +Phased layout requires exact BF16 results; fixed layout has a 1% relative-L2 +budget. These are mechanism tests, with no checkpoint or full-model execution. + +The CLI first runs the small unittest guards, then the integration checks. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import shutil +import sys +import unittest + +import torch + +REPO_ROOT = Path(__file__).resolve().parents[3] +COMPILER_ROOT = REPO_ROOT / "PLENA_Compiler" +for path in (REPO_ROOT / "PLENA_Tools", COMPILER_ROOT, REPO_ROOT): + if str(path) not in sys.path: + sys.path.insert(0, str(path)) + +from compiler.asm_templates._imm import load_large_int # noqa: E402 +from compiler.assembler.assembly_to_binary import AssemblyToBinary # noqa: E402 +from compiler.aten.plena import PlenaCompiler # noqa: E402 +from compiler.aten.plena.isa_matrix_view import ( # noqa: E402 + MatrixViewDescriptor, + MatrixViewMap, + MatrixViewShape, + validate_matrix_view_dominance, +) +from compiler.aten.plena.program_lcompute import ( # noqa: E402 + KIMI_KDA, + NEMOTRON_MAMBA, + MatrixRecurrenceSpec, + MatrixSramPoint, + RecurrenceKind, + RecurrenceLayout, + build_recurrence_field_manifest, + build_recurrence_working_set, + lower_matrix_recurrence, + validate_recurrence_output_stores, +) +from transactional_emulator.testbench.aten import _matrix_lcompute as kit # noqa: E402 +from transactional_emulator.testbench.aten.golden import golden_linear # noqa: E402 +from transactional_emulator.testbench.emulator_runner import ( # noqa: E402 + _parse_matrix_view_packet_counters, + run_and_assert, + run_emulator, +) +from transactional_emulator.testbench.sim_env_utils import create_mem_for_sim # noqa: E402 +from transactional_emulator.tools.create_sim_env import create_sim_env # noqa: E402 + +MLEN, BLEN = 64, 4 +ROWS, K, N = 1, 64, 64 +CONSUMER_WIDTH = 8 +TOKENS = 4 + + +def run_projection(build_dir: Path) -> dict: + torch.manual_seed(20260901) + x = torch.randn(ROWS, K) + x_storage = torch.zeros(BLEN, K) + x_storage[:ROWS].copy_(x) + weight = torch.randn(K, N) + golden = golden_linear(x, weight) + + program = PlenaCompiler(mlen=MLEN, blen=BLEN, mram_tile_capacity=64) + x_input = program.input( + "X", + shape=(ROWS, K), + physical_shape=(BLEN, K), + real_data_ratio=1.0, + ) + w_input = program.input( + "W", + shape=(K, N), + physical_shape=(K, N), + ) + zero_input = program.input( + "zero", + shape=(ROWS, N), + physical_shape=(BLEN, N), + real_data_ratio=1.0, + ) + x_vram = program.load_batch(x_input, name="X_vram") + zero = program.load_batch(zero_input, name="zero_vram") + output_placeholder = program.alloc( + "matrix_output_placeholder", + ROWS, + N, + strict=False, + physical_shape=(BLEN, N), + ) + restored = program.alloc( + "restored", + ROWS, + N, + strict=False, + physical_shape=(BLEN, N), + ) + descriptor = MatrixViewDescriptor( + shape=MatrixViewShape( + rows=ROWS, + cols=CONSUMER_WIDTH, + tile_count=MLEN // CONSUMER_WIDTH, + ), + mapping=MatrixViewMap( + tile_pitch_rows=CONSUMER_WIDTH // BLEN, + ), + ) + matrix_base = program.reserve_matrix_view_scratch_v0("matrix_view_projection") + program.vram_sub_projection_stream_k_accum_to( + x_vram, + 0, + w_input, + 0, + output_placeholder, + 0, + 0, + max_k_tiles=1, + matrix_precision="weights", + set_scale=True, + hbm_element_bytes=1, + matrix_view_descriptor=descriptor, + matrix_view_base=matrix_base, + matrix_view_slot=1, + ) + + gp_dst, gp_matrix, gp_zero = program.register_allocator.allocate_gp(3) + try: + restored_addr = program._compiler.get_vram_addr(restored.name) + zero_addr = program._compiler.get_vram_addr(zero.name) + program._emit( + "\n".join( + [ + *load_large_int(gp_dst, restored_addr), + *load_large_int(gp_matrix, matrix_base), + *load_large_int(gp_zero, zero_addr), + f"V_ADD_VV.MV gp{gp_dst}, gp{gp_matrix}, gp{gp_zero}, 0, 2", + ] + ) + + "\n" + ) + finally: + program.register_allocator.free_gp([gp_dst, gp_matrix, gp_zero]) + + isa = program.compile() + validate_matrix_view_dominance(isa) + assert "L_MVIEW_LOAD" not in isa + assert "L_MVIEW_STORE" not in isa + assert "V_ADD_VV.MV" in isa + + inputs = { + "X": x_storage, + "W": weight, + "zero": torch.zeros(BLEN, N), + } + create_sim_env( + inputs, + isa, + {"original_output": golden}, + [0.0] * 10, + build_dir=str(build_dir), + ) + hbm_addrs = {name: program._compiler.get_hbm_layout(name).hbm_base_addr for name in inputs} + create_mem_for_sim( + data_size=256, + mode="behave_sim", + asm="matrix_view_projection", + data=None, + specified_data_order=["X", "W"], + build_path=build_dir, + input_tensors=inputs, + hbm_addrs=hbm_addrs, + ) + + output_addr = program._compiler.get_vram_addr(restored.name) + (build_dir / "comparison_params.json").write_text( + json.dumps( + { + "start_row_idx": output_addr // MLEN, + "num_rows": ROWS, + "num_batches": ROWS, + "elements_per_batch": N, + "row_dim": MLEN, + }, + indent=2, + ) + ) + (build_dir / "generated_asm_code.asm").write_text(isa) + + dump_names = ( + "mram_dump.bin", + "vram_dump.bin", + "fpsram_dump.bin", + "intsram_dump.bin", + ) + dump_paths = [build_dir / name for name in dump_names] + try: + metrics = run_and_assert( + build_dir, + "matrix-view projection", + mlen=MLEN, + blen=BLEN, + vlen=MLEN, + ) + counters = metrics["matrix_view_packet_counters"] + projection_fragments = N // BLEN + weight_read_packets = projection_fragments * K + producer_write_packets = projection_fragments + consumer_read_packets = 1 + expected_packets = weight_read_packets + producer_write_packets + consumer_read_packets + expected_values = weight_read_packets * MLEN + 2 * N + expected_bank_words = weight_read_packets * (MLEN // BLEN) + producer_write_packets + N // BLEN + assert counters == { + # The physical counter includes ordinary M_MM weight-row reads, + # direct affine accumulator writes, and the restored consumer read. + "packets": expected_packets, + "values": expected_values, + "bank_words": expected_bank_words, + "service_cycles": expected_packets, + "ideal_cycles": expected_packets, + "bank_stall_cycles": 0, + } + return {"case": "projection", "matrix_view_packet_counters": counters} + finally: + for dump_path in dump_paths: + dump_path.unlink(missing_ok=True) + + +def run_recurrence(spec: MatrixRecurrenceSpec, layout: RecurrenceLayout, build_dir: Path) -> dict: + point = MatrixSramPoint() + working_set = build_recurrence_working_set(spec, layout=layout, point=point) + exact = layout is RecurrenceLayout.AFFINE + initial_state = kit._state_seed(spec) + expected_state = initial_state.clone() + make_inputs, reference, packet_values = ( + (kit._mamba_inputs, kit._mamba_reference, kit._mamba_packet_values) + if spec.kind is RecurrenceKind.MAMBA + else (kit._kda_inputs, kit._kda_reference, kit._kda_packet_values) + ) + operands_by_token = tuple(make_inputs(token) for token in range(TOKENS)) + expected_outputs, manifests, assemblies = [], [], [] + field_base = kit._round_up(spec.state_bytes_per_layer, 64) + for operands in operands_by_token: + manifest = build_recurrence_field_manifest(working_set, field_hbm_base=field_base) + assembly = lower_matrix_recurrence( + spec, + layout=layout, + point=point, + state_hbm_base=0, + field_hbm_base=field_base, + ) + validate_recurrence_output_stores(assembly, expected_groups=working_set.groups) + expected_output, expected_state = reference(expected_state, operands) + expected_outputs.append(expected_output) + manifests.append(manifest) + assemblies.append(assembly) + field_base = kit._round_up(manifest.end, 64) + + program = "\n".join(assemblies) + validate_matrix_view_dominance(program) + build_dir.mkdir(parents=True, exist_ok=True) + asm_path = build_dir / "generated_asm_code.asm" + asm_path.write_text(program) + assembler = AssemblyToBinary( + str(COMPILER_ROOT / "doc/operation.svh"), + str(COMPILER_ROOT / "doc/configuration.svh"), + ) + assembler.generate_binary(str(asm_path), str(build_dir / "generated_machine_code.mem")) + + image = bytearray(field_base) + image[: spec.state_bytes_per_layer] = kit._pack_state_hbm(initial_state, working_set) + for operands, manifest in zip(operands_by_token, manifests, strict=True): + for packet in manifest.packets: + kit._write_packet(image, packet, packet_values(packet, operands, working_set)) + (build_dir / "hbm_for_behave_sim.bin").write_bytes(image) + (build_dir / "fp_sram.bin").write_bytes(bytes(64)) + (build_dir / "int_sram.bin").write_bytes(bytes(64)) + settings = kit._write_settings(build_dir, point) + with kit._setting_override(settings): + metrics = run_emulator( + build_dir, + hbm_size=kit._round_up(len(image), 64), + threads=1, + dump_cwd=build_dir, + dump_hbm=True, + ) + + post = (build_dir / "hbm_dump.bin").read_bytes() + actual_state = kit._unpack_state_hbm(post, working_set) + state_error = kit._assert_close("final state", actual_state, expected_state, exact=exact) + output_errors = [] + for token, (manifest, expected) in enumerate(zip(manifests, expected_outputs, strict=True)): + actual_groups = [] + for group in range(working_set.groups): + packet = manifest.packet("output_result", group=group) + values = kit._read_bf16(post, packet.hbm_byte_offset, packet.logical_values) + actual_groups.append(values.reshape(working_set.group_heads, spec.row_elements)) + actual = torch.cat(actual_groups, dim=0) + output_errors.append(kit._assert_close(f"token {token} output", actual, expected, exact=exact)) + for packet in manifest.packets: + if packet.field != "output_result": + start = packet.hbm_byte_offset + end = start + packet.transfer_values * 2 + assert post[start:end] == image[start:end], f"input field overwritten: {packet.key}" + + counters = metrics["matrix_view_packet_counters"] + assert counters["packets"] > 0 + assert counters["service_cycles"] == counters["ideal_cycles"] + counters["bank_stall_cycles"] + if exact: + assert counters["bank_stall_cycles"] == 0 + return { + "case": spec.name, + "layout": layout.value, + "tokens": TOKENS, + "state_error": state_error, + "output_errors": output_errors, + "matrix_view_packet_counters": counters, + } + + +class MatrixLComputeGuards(unittest.TestCase): + def test_head_or_lane_permutations(self): + expected = torch.arange(256, dtype=torch.float32).reshape(4, 64) + for axis in (0, 1): + with self.subTest(axis=axis), self.assertRaisesRegex(AssertionError, "values mismatch"): + kit._assert_close("permutation", expected.roll(1, dims=axis), expected) + + def test_non_finite_values(self): + for bad in (float("nan"), float("inf"), float("-inf")): + for actual, expected in ((torch.tensor([bad]), torch.zeros(1)), (torch.zeros(1), torch.tensor([bad]))): + with self.subTest(bad=bad), self.assertRaisesRegex(AssertionError, "non-finite"): + kit._assert_close("non-finite", actual, expected) + + def test_fixed_budget_and_phased_exactness(self): + expected = torch.linspace(-0.03, 0.03, 256) + kit._assert_close("within budget", expected * 0.992, expected) + with self.assertRaisesRegex(AssertionError, "l2_budget"): + kit._assert_close("outside budget", expected * 0.988, expected) + with self.assertRaisesRegex(AssertionError, "l2_budget"): + kit._assert_close("near zero corruption", torch.full((256,), 2e-7), torch.zeros(256)) + kit._assert_close("identity", torch.tensor([1.0, 0.0]), torch.tensor([1.0, 0.0]), exact=True) + for changed in (torch.tensor([1.0078125, 0.0]), torch.tensor([1.0, -0.0])): + with self.subTest(changed=changed), self.assertRaisesRegex(AssertionError, "exact BF16"): + kit._assert_close("exactness", changed, torch.tensor([1.0, 0.0]), exact=True) + + def test_ten_percent_loss(self): + previous_threads = torch.get_num_threads() + torch.set_num_threads(1) + try: + for spec, inputs, reference in ( + (NEMOTRON_MAMBA, kit._mamba_inputs, kit._mamba_reference), + (KIMI_KDA, kit._kda_inputs, kit._kda_reference), + ): + state = kit._state_seed(spec) + for token in range(TOKENS): + output, state = reference(state, inputs(token)) + with ( + self.subTest(model=spec.name, token=token), + self.assertRaisesRegex(AssertionError, "l2_budget"), + ): + kit._assert_close("gain loss", kit._bf16(output * 0.9), output) + finally: + torch.set_num_threads(previous_threads) + + def test_packet_counters(self): + line = ( + "\x1b[32mINFO\x1b[0m Matrix-view packet counters " + "packets=17 values=128 bank_words=32 service_cycles=19 ideal_cycles=17 bank_stall_cycles=2" + ) + self.assertEqual( + _parse_matrix_view_packet_counters(line), + { + "packets": 17, + "values": 128, + "bank_words": 32, + "service_cycles": 19, + "ideal_cycles": 17, + "bank_stall_cycles": 2, + }, + ) + self.assertIsNone(_parse_matrix_view_packet_counters("Matrix-view packet counters packets=17")) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--case", choices=("projection", "mamba", "kda", "all"), default="all") + parser.add_argument("--build-dir", type=Path, default=Path(__file__).parent / "build/matrix_lcompute") + parser.add_argument("--keep-build", action="store_true") + args = parser.parse_args() + build_root = args.build_dir.resolve() + torch.set_num_threads(1) + guards = unittest.defaultTestLoader.loadTestsFromTestCase(MatrixLComputeGuards) + if not unittest.TextTestRunner().run(guards).wasSuccessful(): + raise SystemExit(1) + results = [] + if args.case in {"projection", "all"}: + build_dir = build_root / "projection" + with kit._setting_override(REPO_ROOT / "plena_settings.toml"): + results.append(run_projection(build_dir)) + if not args.keep_build: + shutil.rmtree(build_dir) + for case, spec in (("mamba", NEMOTRON_MAMBA), ("kda", KIMI_KDA)): + if args.case not in {case, "all"}: + continue + for layout in (RecurrenceLayout.FIXED, RecurrenceLayout.AFFINE): + build_dir = build_root / case / layout.value + result = run_recurrence(spec, layout, build_dir) + results.append(result) + print(json.dumps(result), flush=True) + if not args.keep_build: + shutil.rmtree(build_dir) + print(json.dumps(results, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/transactional_emulator/testbench/emulator_runner.py b/transactional_emulator/testbench/emulator_runner.py index bc329ab6..f03fa443 100644 --- a/transactional_emulator/testbench/emulator_runner.py +++ b/transactional_emulator/testbench/emulator_runner.py @@ -82,6 +82,34 @@ def _run_file_suffix(run_label: str | None) -> str: return f".{safe}" +_MATRIX_VIEW_PACKET_COUNTER_RE = re.compile( + r"Matrix-view packet counters\s+" + r"packets=(?P\d+)\s+" + r"values=(?P\d+)\s+" + r"bank_words=(?P\d+)\s+" + r"service_cycles=(?P\d+)\s+" + r"ideal_cycles=(?P\d+)\s+" + r"bank_stall_cycles=(?P\d+)" +) + +_ANSI_ESCAPE_RE = re.compile(r"\x1b(?:\[[0-?]*[ -/]*[@-~]|\][^\x07]*(?:\x07|\x1b\\))") + + +def _strip_ansi(value: str) -> str: + """Remove terminal colour/control sequences before parsing Rust logs.""" + + return _ANSI_ESCAPE_RE.sub("", value) + + +def _parse_matrix_view_packet_counters(line: str) -> dict[str, int] | None: + """Parse Matrix-SRAM view counters without depending on log prefixes.""" + + match = _MATRIX_VIEW_PACKET_COUNTER_RE.search(_strip_ansi(line)) + if match is None: + return None + return {name: int(value) for name, value in match.groupdict().items()} + + def run_emulator( build_dir: Path, hbm_size: int | None = None, @@ -91,6 +119,7 @@ def run_emulator( run_label: str | None = None, timing_model: str | None = None, dump_cwd: Path | None = None, + dump_hbm: bool = False, ) -> dict: """Run the Rust transactional emulator with build artifacts from build_dir. @@ -123,10 +152,14 @@ def run_emulator( the historical emulator directory. Parallel replay can set this to build_dir so vram_dump.bin/fpsram_dump.bin are not shared between concurrent emulator processes. + dump_hbm: explicitly retain the post-run HBM image in ``build_dir``. + This is intended for connected state/output numerical tests; + ordinary tests leave it disabled because model HBM can be huge. """ + build_dir = Path(build_dir).resolve() emulator_dir = Path(__file__).parent.parent # transactional_emulator/ binary = emulator_dir / "target" / "release" / "transactional_emulator" - dump_dir = Path(dump_cwd) if dump_cwd is not None else emulator_dir + dump_dir = Path(dump_cwd).resolve() if dump_cwd is not None else emulator_dir dump_dir.mkdir(parents=True, exist_ok=True) if stage_profile is None: @@ -210,6 +243,8 @@ def run_emulator( ] if timing_model != "serial": cmd += ["--timing-model", timing_model] + if dump_hbm: + cmd += ["--hbm-dump", str((build_dir / "hbm_dump.bin").resolve())] # tch's download-libtorch stores libtorch in the Cargo build cache. # The binary needs LD_LIBRARY_PATH to find it at runtime. @@ -254,6 +289,7 @@ def run_emulator( "log_path": str(log_path), "stage_profile_requested": bool(stage_profile), "timing_model": timing_model, + "hbm_dump_requested": dump_hbm, } if run_label: metrics["run_label"] = run_label @@ -289,7 +325,9 @@ def run_emulator( print(line, end="") log_file.write(line) - sim_match = sim_latency_re.search(line) + parsed_line = _strip_ansi(line) + + sim_match = sim_latency_re.search(parsed_line) if sim_match: sim_latency_ns = float(sim_match.group(1)) metrics["sim_latency_ns"] = sim_latency_ns @@ -297,18 +335,22 @@ def run_emulator( if sim_match.group(2) is not None: metrics["sim_latency_cycles"] = int(sim_match.group(2)) - topo_match = topology_re.search(line) + topo_match = topology_re.search(parsed_line) if topo_match: metrics["emu_mlen"] = int(topo_match.group(1)) metrics["emu_vlen"] = int(topo_match.group(2)) metrics["emu_blen"] = int(topo_match.group(3)) - hbm_match = hbm_stats_re.search(line) + hbm_match = hbm_stats_re.search(parsed_line) if hbm_match: metrics["hbm_bytes_read"] = int(hbm_match.group(1)) metrics["hbm_bytes_written"] = int(hbm_match.group(2)) metrics["hbm_utilization_bytes_per_sec"] = float(hbm_match.group(3)) + matrix_view_counters = _parse_matrix_view_packet_counters(parsed_line) + if matrix_view_counters is not None: + metrics["matrix_view_packet_counters"] = matrix_view_counters + return_code = proc.wait() ended_at = datetime.now(UTC)