#[cfg(feature = "gpart_adapter")]
pub const GPART_MAGIC: &[u8; 5] = b"GPART";
#[cfg(feature = "gpart_adapter")]
pub const GPART_VERSION: u32 = 1;
#[cfg(feature = "gpart_adapter")]
#[derive(Clone, Debug)]
pub struct GpartAdapter {
pub d: usize,
pub seed: u64,
pub theta: Vec<f32>,
}
#[cfg(feature = "gpart_adapter")]
impl GpartAdapter {
pub fn apply(&self, base_weights: &mut [f32]) {
let n = base_weights.len();
if n == 0 || self.d == 0 {
return;
}
let assignments = self.generate_assignments(n);
let group_sizes = self.compute_group_sizes(n, &assignments);
let mut group_delta = vec![0.0f32; self.d];
for g in 0..self.d {
let scale = 1.0 / (group_sizes[g] as f32).sqrt();
group_delta[g] = scale * self.theta[g];
}
for i in 0..n {
base_weights[i] += group_delta[assignments[i]];
}
}
pub fn apply_with_scratch(
&self,
base_weights: &mut [f32],
assignments: &mut [usize],
group_sizes: &mut [usize],
) {
let n = base_weights.len();
if n == 0 || self.d == 0 {
return;
}
self.generate_assignments_into(n, assignments);
self.compute_group_sizes_into(n, assignments, group_sizes);
let mut group_delta = vec![0.0f32; self.d];
for g in 0..self.d {
let scale = 1.0 / (group_sizes[g] as f32).sqrt();
group_delta[g] = scale * self.theta[g];
}
for i in 0..n {
base_weights[i] += group_delta[assignments[i]];
}
}
#[cfg(feature = "gpart_pruning")]
pub fn apply_with_scratch_masked(
&self,
base_weights: &mut [f32],
assignments: &mut [usize],
group_sizes: &mut [usize],
group_mask: &[bool],
) {
let n = base_weights.len();
if n == 0 || self.d == 0 {
return;
}
debug_assert!(
group_mask.len() >= self.d,
"group_mask len {} < d {}",
group_mask.len(),
self.d,
);
self.generate_assignments_into(n, assignments);
self.compute_group_sizes_into(n, assignments, group_sizes);
let mut group_delta: Vec<f32> = Vec::with_capacity(self.d);
for g in 0..self.d {
let active = group_mask[g];
let scale = 1.0 / (group_sizes[g] as f32).sqrt();
group_delta.push(scale * self.theta[g] * (active as u8 as f32));
}
for i in 0..n {
base_weights[i] += group_delta[assignments[i]];
}
}
pub fn apply_simd(&self, base_weights: &mut [f32]) {
let n = base_weights.len();
if n == 0 || self.d == 0 {
return;
}
let assignments = self.generate_assignments(n);
let group_sizes = self.compute_group_sizes(n, &assignments);
let group_delta: Vec<f32> = (0..self.d)
.map(|g| self.theta[g] / (group_sizes[g] as f32).sqrt())
.collect();
let chunks = n / 8;
for c in 0..chunks {
let base = c * 8;
for j in 0..8 {
let i = base + j;
base_weights[i] += group_delta[assignments[i]];
}
}
for i in (chunks * 8)..n {
base_weights[i] += group_delta[assignments[i]];
}
}
#[cfg(feature = "gpart_pruning")]
pub fn topk_mask(&self, k: usize) -> Vec<bool> {
if k >= self.d {
return vec![true; self.d];
}
if k == 0 {
return vec![false; self.d];
}
let mut magnitudes: Vec<(f32, usize)> =
(0..self.d).map(|g| (self.theta[g].abs(), g)).collect();
magnitudes.select_nth_unstable_by(k - 1, |a, b| b.0.total_cmp(&a.0));
let mut mask = vec![false; self.d];
for &(_, g) in &magnitudes[..k] {
mask[g] = true;
}
mask
}
fn generate_assignments(&self, n: usize) -> Vec<usize> {
let mut rng = fastrand::Rng::with_seed(self.seed);
let mut assignments = Vec::with_capacity(n);
for _ in 0..n {
assignments.push(rng.usize(..self.d));
}
assignments
}
fn generate_assignments_into(&self, n: usize, out: &mut [usize]) {
let mut rng = fastrand::Rng::with_seed(self.seed);
for i in 0..n.min(out.len()) {
out[i] = rng.usize(..self.d);
}
}
fn compute_group_sizes(&self, n: usize, assignments: &[usize]) -> Vec<usize> {
let mut sizes = vec![0usize; self.d];
for &g in &assignments[..n] {
sizes[g] += 1;
}
sizes
}
fn compute_group_sizes_into(&self, n: usize, assignments: &[usize], out: &mut [usize]) {
for g in out.iter_mut().take(self.d) {
*g = 0;
}
for &g in &assignments[..n] {
out[g] += 1;
}
}
pub fn commitment(&self) -> [u8; 32] {
let mut hasher = blake3::Hasher::new();
hasher.update(&self.seed.to_le_bytes());
let theta_bytes: &[u8] = unsafe {
std::slice::from_raw_parts(
self.theta.as_ptr() as *const u8,
self.theta.len() * std::mem::size_of::<f32>(),
)
};
hasher.update(theta_bytes);
*hasher.finalize().as_bytes()
}
pub fn verify(&self, expected: &[u8; 32]) -> bool {
self.commitment() == *expected
}
pub fn save(&self, path: &std::path::Path) -> Result<(), String> {
let d_bytes = self.d as u32;
let theta_byte_len = self.theta.len() * std::mem::size_of::<f32>();
let total = 5 + 4 + 4 + 8 + 32 + theta_byte_len;
let mut buf = Vec::with_capacity(total);
buf.extend_from_slice(GPART_MAGIC);
buf.extend_from_slice(&GPART_VERSION.to_le_bytes());
buf.extend_from_slice(&d_bytes.to_le_bytes());
buf.extend_from_slice(&self.seed.to_le_bytes());
let commit_offset = buf.len();
buf.extend_from_slice(&[0u8; 32]);
let theta_bytes: &[u8] =
unsafe { std::slice::from_raw_parts(self.theta.as_ptr() as *const u8, theta_byte_len) };
buf.extend_from_slice(theta_bytes);
let commitment = self.commitment();
buf[commit_offset..commit_offset + 32].copy_from_slice(&commitment);
std::fs::write(path, &buf).map_err(|e| format!("Failed to write gpart file: {e}"))
}
pub fn load(path: &std::path::Path) -> Result<Self, String> {
let file_data =
std::fs::read(path).map_err(|e| format!("Failed to read gpart file: {e}"))?;
if file_data.len() < 53 {
return Err("File too small for gpart header".into());
}
if &file_data[0..5] != GPART_MAGIC {
return Err("Invalid gpart magic bytes".into());
}
let version = u32::from_le_bytes(
file_data[5..9]
.try_into()
.map_err(|e: std::array::TryFromSliceError| format!("Version parse: {e}"))?,
);
if version != GPART_VERSION {
return Err(format!("Unsupported gpart version: {version}"));
}
let d = u32::from_le_bytes(
file_data[9..13]
.try_into()
.map_err(|e: std::array::TryFromSliceError| format!("d parse: {e}"))?,
) as usize;
let seed = u64::from_le_bytes(
file_data[13..21]
.try_into()
.map_err(|e: std::array::TryFromSliceError| format!("seed parse: {e}"))?,
);
let stored_commitment = &file_data[21..53];
let theta_bytes_start = 53;
let theta_bytes_len = d * std::mem::size_of::<f32>();
if theta_bytes_start + theta_bytes_len > file_data.len() {
return Err("Truncated theta data".into());
}
let theta: Vec<f32> = {
#[cfg(target_endian = "little")]
{
let src = &file_data[theta_bytes_start..theta_bytes_start + theta_bytes_len];
let count = d;
let mut v = Vec::with_capacity(count);
unsafe {
std::ptr::copy_nonoverlapping(
src.as_ptr(),
v.as_mut_ptr() as *mut u8,
theta_bytes_len,
);
v.set_len(count);
}
v
}
#[cfg(not(target_endian = "little"))]
{
file_data[theta_bytes_start..theta_bytes_start + theta_bytes_len]
.chunks_exact(4)
.map(|c| f32::from_le_bytes(c.try_into().expect("chunk is 4 bytes")))
.collect()
}
};
let adapter = Self { d, seed, theta };
let computed = adapter.commitment();
if computed != stored_commitment {
return Err("GPart file commitment mismatch".into());
}
Ok(adapter)
}
pub fn storage_bytes(&self) -> usize {
8 + self.theta.len() * std::mem::size_of::<f32>()
}
pub fn check_isometry(&self, n: usize) -> bool {
if n == 0 || self.d == 0 {
return true;
}
let assignments = self.generate_assignments(n);
let group_sizes = self.compute_group_sizes(n, &assignments);
let mut group_delta = vec![0.0f32; self.d];
for g in 0..self.d {
let scale = 1.0 / (group_sizes[g] as f32).sqrt();
group_delta[g] = scale * self.theta[g];
}
let mut projected_norm_sq = 0.0f32;
for &g in assignments.iter().take(n) {
let delta = group_delta[g];
projected_norm_sq += delta * delta;
}
let theta_norm_sq: f32 = self.theta.iter().map(|&v| v * v).sum();
(projected_norm_sq - theta_norm_sq).abs() < 1e-3
}
}
#[cfg(feature = "gpart_adapter")]
#[derive(Clone, Debug)]
pub struct GpartPrepared {
deltas: Vec<f32>,
}
#[cfg(feature = "gpart_adapter")]
impl GpartAdapter {
pub fn prepare(&self, n: usize) -> GpartPrepared {
if n == 0 || self.d == 0 {
return GpartPrepared { deltas: Vec::new() };
}
let assignments = self.generate_assignments(n);
let group_sizes = self.compute_group_sizes(n, &assignments);
let group_delta: Vec<f32> = (0..self.d)
.map(|g| {
let scale = 1.0 / (group_sizes[g] as f32).sqrt();
scale * self.theta[g]
})
.collect();
let deltas = assignments.iter().map(|&g| group_delta[g]).collect();
GpartPrepared { deltas }
}
}
#[cfg(feature = "gpart_adapter")]
impl GpartPrepared {
pub fn apply(&self, base_weights: &mut [f32]) {
let len = base_weights.len().min(self.deltas.len());
for (w, &delta) in base_weights.iter_mut().zip(self.deltas.iter()).take(len) {
*w += delta;
}
}
}
#[cfg(feature = "gpart_adapter")]
#[derive(Clone, Debug)]
pub struct GpartPair {
pub reader: Option<GpartAdapter>,
pub writer: Option<GpartAdapter>,
}
#[cfg(feature = "gpart_adapter")]
impl GpartPair {
pub fn none() -> Self {
Self {
reader: None,
writer: None,
}
}
pub fn apply_prefill(&self, base_weights: &mut [f32]) {
if let Some(ref adapter) = self.reader {
adapter.apply(base_weights);
}
}
pub fn apply_decode(&self, base_weights: &mut [f32]) {
if let Some(ref adapter) = self.writer {
adapter.apply(base_weights);
}
}
}
#[cfg(feature = "gpart_adapter")]
impl TryFrom<&crate::LoraAdapter> for GpartAdapter {
type Error = &'static str;
fn try_from(_lora: &crate::LoraAdapter) -> Result<Self, Self::Error> {
Err("GpartAdapter requires pre-computed θ_d from riir-ai training pipeline (P⁺ΔW)")
}
}
#[cfg(feature = "gpart_adapter")]
#[allow(clippy::eq_op, clippy::assertions_on_constants)]
const _: () = assert!(8 + 90 * 4 <= 368);