Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 13 additions & 1 deletion .github/workflows/transactional_emulator.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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'
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
3 changes: 3 additions & 0 deletions justfile
Original file line number Diff line number Diff line change
Expand Up @@ -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}}
9 changes: 9 additions & 0 deletions plena_settings.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
26 changes: 26 additions & 0 deletions transactional_emulator/lib/quantize/src/dtype.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand All @@ -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;

Expand Down Expand Up @@ -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).
Expand Down
8 changes: 0 additions & 8 deletions transactional_emulator/lib/sram/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,14 +48,6 @@ impl<T> Cell<T> {
}
}

impl Cell<QuantTensor> {
/// 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 {
Expand Down
Loading
Loading