use anyhow::Result;
use candle_core::{DType, Device, Tensor};
use candle_nn::{Init, VarMap};
use finetype_model::sibling_context::SiblingContextConfig;
pub struct SiblingContextTrainable {
blocks: Vec<TrainableTransformerBlock>,
final_norm: TrainableLayerNorm,
embed_dim: usize,
}
struct TrainableTransformerBlock {
norm1: TrainableLayerNorm,
attn: TrainableMultiHeadAttention,
norm2: TrainableLayerNorm,
ffn: TrainableFFN,
}
struct TrainableLayerNorm {
weight: Tensor, bias: Tensor, }
struct TrainableMultiHeadAttention {
wq: Tensor, bq: Tensor, wk: Tensor, bk: Tensor, wv: Tensor, bv: Tensor, out_weight: Tensor, out_bias: Tensor, n_heads: usize,
head_dim: usize,
}
struct TrainableFFN {
w1: Tensor, b1: Tensor, w2: Tensor, b2: Tensor, }
impl SiblingContextTrainable {
pub fn new(varmap: &VarMap, config: &SiblingContextConfig, device: &Device) -> Result<Self> {
let d = config.embed_dim;
let ff = d * 4;
let n_heads = config.n_heads;
let head_dim = d / n_heads;
let scale = (1.0f64 / d as f64).sqrt();
let mut blocks = Vec::with_capacity(config.n_layers);
for i in 0..config.n_layers {
let prefix = format!("blocks.{}", i);
let norm1 = TrainableLayerNorm {
weight: varmap.get(
(d,),
&format!("{prefix}.norm1.weight"),
Init::Const(1.0),
DType::F32,
device,
)?,
bias: varmap.get(
(d,),
&format!("{prefix}.norm1.bias"),
Init::Const(0.0),
DType::F32,
device,
)?,
};
let attn = TrainableMultiHeadAttention {
wq: varmap.get(
(d, d),
&format!("{prefix}.attn.wq"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
bq: varmap.get(
(d,),
&format!("{prefix}.attn.bq"),
Init::Const(0.0),
DType::F32,
device,
)?,
wk: varmap.get(
(d, d),
&format!("{prefix}.attn.wk"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
bk: varmap.get(
(d,),
&format!("{prefix}.attn.bk"),
Init::Const(0.0),
DType::F32,
device,
)?,
wv: varmap.get(
(d, d),
&format!("{prefix}.attn.wv"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
bv: varmap.get(
(d,),
&format!("{prefix}.attn.bv"),
Init::Const(0.0),
DType::F32,
device,
)?,
out_weight: varmap.get(
(d, d),
&format!("{prefix}.attn.out_weight"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
out_bias: varmap.get(
(d,),
&format!("{prefix}.attn.out_bias"),
Init::Const(0.0),
DType::F32,
device,
)?,
n_heads,
head_dim,
};
let norm2 = TrainableLayerNorm {
weight: varmap.get(
(d,),
&format!("{prefix}.norm2.weight"),
Init::Const(1.0),
DType::F32,
device,
)?,
bias: varmap.get(
(d,),
&format!("{prefix}.norm2.bias"),
Init::Const(0.0),
DType::F32,
device,
)?,
};
let ffn = TrainableFFN {
w1: varmap.get(
(ff, d),
&format!("{prefix}.ffn.w1"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
b1: varmap.get(
(ff,),
&format!("{prefix}.ffn.b1"),
Init::Const(0.0),
DType::F32,
device,
)?,
w2: varmap.get(
(d, ff),
&format!("{prefix}.ffn.w2"),
Init::Randn {
mean: 0.0,
stdev: scale,
},
DType::F32,
device,
)?,
b2: varmap.get(
(d,),
&format!("{prefix}.ffn.b2"),
Init::Const(0.0),
DType::F32,
device,
)?,
};
blocks.push(TrainableTransformerBlock {
norm1,
attn,
norm2,
ffn,
});
}
let final_norm = TrainableLayerNorm {
weight: varmap.get(
(d,),
"final_norm.weight",
Init::Const(1.0),
DType::F32,
device,
)?,
bias: varmap.get(
(d,),
"final_norm.bias",
Init::Const(0.0),
DType::F32,
device,
)?,
};
Ok(Self {
blocks,
final_norm,
embed_dim: d,
})
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let mut out = x.clone();
for block in &self.blocks {
out = block.forward(&out)?;
}
self.final_norm.forward(&out)
}
pub fn param_count(&self) -> usize {
let d = self.embed_dim;
let ff = d * 4;
let per_block = 4 * (d * d + d) + 2 * (d + d) + (ff * d + ff) + (d * ff + d); let final_norm = d + d;
per_block * self.blocks.len() + final_norm
}
}
impl TrainableTransformerBlock {
fn forward(&self, x: &Tensor) -> Result<Tensor> {
let normed = self.norm1.forward(x)?;
let attn_out = self.attn.forward(&normed)?;
let x = (x + &attn_out)?;
let normed = self.norm2.forward(&x)?;
let ffn_out = self.ffn.forward(&normed)?;
Ok((&x + &ffn_out)?)
}
}
impl TrainableLayerNorm {
fn forward(&self, x: &Tensor) -> Result<Tensor> {
let eps = 1e-5_f64;
let d = x.dim(1)?;
let mean = (x.sum(1)? / d as f64)?;
let mean = mean.unsqueeze(1)?;
let diff = x.broadcast_sub(&mean)?;
let var = ((&diff * &diff)?.sum(1)? / d as f64)?;
let std = (var + eps)?.sqrt()?.unsqueeze(1)?;
let normed = diff.broadcast_div(&std)?;
Ok(normed
.broadcast_mul(&self.weight)?
.broadcast_add(&self.bias)?)
}
}
impl TrainableMultiHeadAttention {
fn forward(&self, x: &Tensor) -> Result<Tensor> {
let n = x.dim(0)?;
let d = x.dim(1)?;
let h = self.n_heads;
let hd = self.head_dim;
let q = x.matmul(&self.wq.t()?)?.broadcast_add(&self.bq)?;
let k = x.matmul(&self.wk.t()?)?.broadcast_add(&self.bk)?;
let v = x.matmul(&self.wv.t()?)?.broadcast_add(&self.bv)?;
let q = q.reshape((n, h, hd))?.transpose(0, 1)?;
let k = k.reshape((n, h, hd))?.transpose(0, 1)?;
let v = v.reshape((n, h, hd))?.transpose(0, 1)?;
let scale = (hd as f64).sqrt();
let attn_weights = (q.matmul(&k.transpose(1, 2)?)? / scale)?;
let attn_max = attn_weights.max(2)?.unsqueeze(2)?;
let shifted = attn_weights.broadcast_sub(&attn_max)?;
let exp = shifted.exp()?;
let sum_exp = exp.sum(2)?.unsqueeze(2)?;
let attn_probs = exp.broadcast_div(&sum_exp)?;
let attn_out = attn_probs.matmul(&v)?;
let attn_out = attn_out.transpose(0, 1)?.reshape((n, d))?;
Ok(attn_out
.matmul(&self.out_weight.t()?)?
.broadcast_add(&self.out_bias)?)
}
}
impl TrainableFFN {
fn forward(&self, x: &Tensor) -> Result<Tensor> {
let h = x.matmul(&self.w1.t()?)?.broadcast_add(&self.b1)?;
let h = h.gelu_erf()?;
Ok(h.matmul(&self.w2.t()?)?.broadcast_add(&self.b2)?)
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_nn::{AdamW, Optimizer, ParamsAdamW};
#[test]
fn test_trainable_param_count() {
let varmap = VarMap::new();
let config = SiblingContextConfig::default();
let device = Device::Cpu;
let model = SiblingContextTrainable::new(&varmap, &config, &device).unwrap();
assert_eq!(model.param_count(), 396800);
let varmap_params: usize = varmap
.all_vars()
.iter()
.map(|v| v.as_tensor().elem_count())
.sum();
assert_eq!(varmap_params, 396800);
}
#[test]
fn test_trainable_forward_shape() {
let varmap = VarMap::new();
let config = SiblingContextConfig::default();
let device = Device::Cpu;
let model = SiblingContextTrainable::new(&varmap, &config, &device).unwrap();
for n in [1, 5, 10, 20] {
let input = Tensor::randn(0.0f32, 1.0, (n, 128), &device).unwrap();
let output = model.forward(&input).unwrap();
assert_eq!(output.dims(), &[n, 128], "Shape mismatch for N={}", n);
}
}
#[test]
fn test_trainable_save_load_round_trip() {
let varmap = VarMap::new();
let config = SiblingContextConfig::default();
let device = Device::Cpu;
let model = SiblingContextTrainable::new(&varmap, &config, &device).unwrap();
let tmp_dir = std::env::temp_dir().join("finetype_sibling_train_test");
let _ = std::fs::remove_dir_all(&tmp_dir);
std::fs::create_dir_all(&tmp_dir).unwrap();
let model_path = tmp_dir.join("model.safetensors");
varmap.save(&model_path).unwrap();
let config_json = serde_json::to_string_pretty(&config).unwrap();
std::fs::write(tmp_dir.join("config.json"), &config_json).unwrap();
let loaded = finetype_model::SiblingContextAttention::load(&tmp_dir).unwrap();
assert_eq!(loaded.param_count(), 396800);
let input = Tensor::randn(0.0f32, 1.0, (5, 128), &device).unwrap();
let out_train: Vec<f32> = model
.forward(&input)
.unwrap()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
let out_infer: Vec<f32> = loaded
.forward(&input)
.unwrap()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
for (a, b) in out_train.iter().zip(out_infer.iter()) {
assert!((a - b).abs() < 1e-5, "Round-trip mismatch: {} vs {}", a, b);
}
let _ = std::fs::remove_dir_all(&tmp_dir);
}
#[test]
fn test_gradient_flow() {
let varmap = VarMap::new();
let config = SiblingContextConfig::default();
let device = Device::Cpu;
let model = SiblingContextTrainable::new(&varmap, &config, &device).unwrap();
let initial: Vec<f32> = varmap.all_vars()[0]
.as_tensor()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
let input = Tensor::randn(0.0f32, 1.0, (3, 128), &device).unwrap();
let target = Tensor::randn(0.0f32, 1.0, (3, 128), &device).unwrap();
let output = model.forward(&input).unwrap();
let loss = (&output - &target)
.unwrap()
.sqr()
.unwrap()
.mean_all()
.unwrap();
let adamw_params = ParamsAdamW {
lr: 1e-2,
weight_decay: 0.0,
..Default::default()
};
let mut optimizer = AdamW::new(varmap.all_vars(), adamw_params).unwrap();
optimizer.backward_step(&loss).unwrap();
let updated: Vec<f32> = varmap.all_vars()[0]
.as_tensor()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
let max_diff: f32 = initial
.iter()
.zip(updated.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_diff > 1e-8,
"Weights should change after backward_step, max_diff={}",
max_diff
);
}
}