use crate::int8::{Int8Tier, QuantizedMatrix, dot_i32};
use std::sync::{Condvar, Mutex, OnceLock};
#[derive(Clone, Copy)]
struct Job {
x_q: *const i8,
x_scales: *const f32,
w_data: *const i8,
w_scales: *const f32,
bias: *const f32,
out: *mut f32,
m: usize,
n: usize,
k: usize,
tier: Int8Tier,
partitions: usize,
}
unsafe impl Send for Job {}
unsafe impl Sync for Job {}
struct Control {
generation: u64,
job: Option<Job>,
remaining: usize,
}
struct Shared {
control: Mutex<Control>,
go: Condvar,
done: Condvar,
}
pub struct Team {
shared: &'static Shared,
partitions: usize,
dispatch_gate: Mutex<()>,
}
pub fn armed() -> Option<&'static Team> {
static TEAM: OnceLock<Option<Team>> = OnceLock::new();
TEAM.get_or_init(|| {
let ceiling = std::thread::available_parallelism().map_or(1, usize::from);
let requested: usize = std::env::var("FTTS_INT8_THREADS")
.ok()
.and_then(|value| value.parse().ok())
.unwrap_or(6);
let partitions = requested.min(ceiling);
if partitions <= 1 {
return None;
}
let shared: &'static Shared = Box::leak(Box::new(Shared {
control: Mutex::new(Control {
generation: 0,
job: None,
remaining: 0,
}),
go: Condvar::new(),
done: Condvar::new(),
}));
for worker in 1..partitions {
std::thread::Builder::new()
.name(format!("ftts-int8-{worker}"))
.spawn(move || worker_loop(shared, worker))
.expect("spawn int8 worker");
}
Some(Team {
shared,
partitions,
dispatch_gate: Mutex::new(()),
})
})
.as_ref()
}
fn worker_loop(shared: &'static Shared, worker: usize) {
let mut seen = 0_u64;
loop {
let job = {
let mut control = shared.control.lock().expect("team control poisoned");
while control.generation == seen {
control = shared.go.wait(control).expect("team control poisoned");
}
seen = control.generation;
control.job.expect("generation bumped without a job")
};
run_partition(&job, worker);
let mut control = shared.control.lock().expect("team control poisoned");
control.remaining -= 1;
if control.remaining == 0 {
shared.done.notify_all();
}
}
}
fn run_partition(job: &Job, worker: usize) {
let chunk = job.n.div_ceil(job.partitions);
let start = (worker * chunk).min(job.n);
let end = ((worker + 1) * chunk).min(job.n);
if start >= end {
return;
}
let (x_q, x_scales, w_data, w_scales, bias, out) = unsafe {
(
std::slice::from_raw_parts(job.x_q, job.m * job.k),
std::slice::from_raw_parts(job.x_scales, job.m),
std::slice::from_raw_parts(job.w_data, job.n * job.k),
std::slice::from_raw_parts(job.w_scales, job.n),
(!job.bias.is_null()).then(|| std::slice::from_raw_parts(job.bias, job.n)),
std::slice::from_raw_parts_mut(job.out, job.m * job.n),
)
};
for col in start..end {
let w_row = &w_data[col * job.k..(col + 1) * job.k];
let w_scale = w_scales[col];
let bias_term = bias.map(|b| b[col]);
for row in 0..job.m {
let x_row = &x_q[row * job.k..(row + 1) * job.k];
let acc = dot_i32(x_row, w_row, job.tier);
let value = acc as f32 * (x_scales[row] * w_scale);
out[row * job.n + col] = bias_term.map_or(value, |b| value + b);
}
}
}
impl Team {
#[allow(clippy::too_many_arguments)]
pub fn linear_q8(
&self,
x_q: &[i8],
x_scales: &[f32],
weight: &QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
tier: Int8Tier,
) {
let (n, k) = (weight.n, weight.k);
assert_eq!(x_q.len(), m * k, "x_q must be [m, k]");
assert_eq!(x_scales.len(), m, "x_scales must be [m]");
assert_eq!(out.len(), m * n, "out must be [m, n]");
if let Some(bias) = bias {
assert_eq!(bias.len(), n, "bias must be [n]");
}
let job = Job {
x_q: x_q.as_ptr(),
x_scales: x_scales.as_ptr(),
w_data: weight.data.as_ptr(),
w_scales: weight.scales.as_ptr(),
bias: bias.map_or(std::ptr::null(), <[f32]>::as_ptr),
out: out.as_mut_ptr(),
m,
n,
k,
tier,
partitions: self.partitions,
};
let _gate = self.dispatch_gate.lock().expect("dispatch gate poisoned");
{
let mut control = self.shared.control.lock().expect("team control poisoned");
control.job = Some(job);
control.generation += 1;
control.remaining = self.partitions - 1;
self.shared.go.notify_all();
}
run_partition(&job, 0);
let mut control = self.shared.control.lock().expect("team control poisoned");
while control.remaining > 0 {
control = self
.shared
.done
.wait(control)
.expect("team control poisoned");
}
control.job = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::int8::linear_q8;
fn matrix(n: usize, k: usize, seed: u64) -> QuantizedMatrix {
let mut state = seed;
let data: Vec<i8> = (0..n * k)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
(((state >> 33) % 255) as i32 - 127) as i8
})
.collect();
let scales: Vec<f32> = (0..n).map(|row| 0.001 + (row % 7) as f32 * 0.01).collect();
QuantizedMatrix { data, scales, n, k }
}
fn test_team(partitions: usize) -> Team {
let shared: &'static Shared = Box::leak(Box::new(Shared {
control: Mutex::new(Control {
generation: 0,
job: None,
remaining: 0,
}),
go: Condvar::new(),
done: Condvar::new(),
}));
for worker in 1..partitions {
std::thread::spawn(move || worker_loop(shared, worker));
}
Team {
shared,
partitions,
dispatch_gate: Mutex::new(()),
}
}
#[test]
fn every_partition_count_is_bit_identical_to_serial_at_model_shapes() {
for &(m, n, k) in &[
(1_usize, 2048_usize, 1024_usize),
(1, 1024, 3072),
(16, 3072, 1024),
(2, 517, 129), ] {
let weight = matrix(n, k, 42 ^ (n as u64) << 20);
let x_q: Vec<i8> = (0..m * k).map(|i| ((i * 31 + 7) % 255) as i8).collect();
let x_scales: Vec<f32> = (0..m).map(|row| 0.02 + row as f32 * 0.005).collect();
let mut serial = vec![0.0_f32; m * n];
linear_q8(
&x_q,
&x_scales,
&weight,
None,
m,
&mut serial,
Int8Tier::Scalar,
);
for partitions in [2_usize, 3, 4, 8] {
let team = test_team(partitions);
let mut parallel = vec![0.0_f32; m * n];
team.linear_q8(
&x_q,
&x_scales,
&weight,
None,
m,
&mut parallel,
Int8Tier::Scalar,
);
for (index, (a, b)) in serial.iter().zip(¶llel).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"partitions={partitions} m={m} n={n} k={k} element {index}"
);
}
}
}
}
#[test]
fn thousands_of_mixed_dispatches_complete_without_deadlock() {
let team = test_team(4);
let weight_a = matrix(256, 512, 7);
let weight_b = matrix(96, 128, 11);
let x_a: Vec<i8> = vec![3; 512];
let x_b: Vec<i8> = vec![-5; 2 * 128];
let mut out_a = vec![0.0_f32; 256];
let mut out_b = vec![0.0_f32; 2 * 96];
for _ in 0..2_000 {
team.linear_q8(
&x_a,
&[0.5],
&weight_a,
None,
1,
&mut out_a,
Int8Tier::Scalar,
);
team.linear_q8(
&x_b,
&[0.5, 0.25],
&weight_b,
None,
2,
&mut out_b,
Int8Tier::Scalar,
);
}
assert!(out_a.iter().all(|value| value.is_finite()));
assert!(out_b.iter().all(|value| value.is_finite()));
}
#[test]
fn bias_reaches_every_partition() {
let (m, n, k) = (2_usize, 130_usize, 64_usize);
let weight = matrix(n, k, 99);
let bias: Vec<f32> = (0..n).map(|i| i as f32).collect();
let x_q: Vec<i8> = vec![1; m * k];
let x_scales = vec![1.0_f32; m];
let mut serial = vec![0.0_f32; m * n];
linear_q8(
&x_q,
&x_scales,
&weight,
Some(&bias),
m,
&mut serial,
Int8Tier::Scalar,
);
let team = test_team(3);
let mut parallel = vec![0.0_f32; m * n];
team.linear_q8(
&x_q,
&x_scales,
&weight,
Some(&bias),
m,
&mut parallel,
Int8Tier::Scalar,
);
assert_eq!(
serial.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
parallel.iter().map(|v| v.to_bits()).collect::<Vec<_>>()
);
}
}