use crate::candle::candle_core::{Device, Tensor};
use crate::candle::candle_nn::VarMap;
pub const NEW_TOPIC_LOGIT_BIAS: f64 = -10.0;
#[derive(Clone, Copy, Debug, Default)]
pub struct Growth {
pub add_topics: usize,
pub add_embedding_dim: usize,
}
impl Growth {
#[must_use]
pub fn is_none(self) -> bool {
self.add_topics == 0 && self.add_embedding_dim == 0
}
}
pub struct AxisRemap<'a> {
pub new_to_old: &'a [Option<usize>],
pub n_old: usize,
}
pub struct GrowthDims<'a> {
pub k_old: usize,
pub k_new: usize,
pub h_old: usize,
pub h_new: usize,
pub gene_axis: Option<AxisRemap<'a>>,
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Axis {
Topics,
Embedding,
Features,
}
impl Axis {
#[must_use]
pub fn label(self) -> &'static str {
match self {
Axis::Topics => "K",
Axis::Embedding => "H",
Axis::Features => "D",
}
}
}
impl GrowthDims<'_> {
pub fn classify(&self, old: usize, new: usize) -> Option<Axis> {
if self.is_gene_axis(old, new) {
Some(Axis::Features)
} else if old == self.k_old && new == self.k_new {
Some(Axis::Topics)
} else if old == self.h_old && new == self.h_new {
Some(Axis::Embedding)
} else {
None
}
}
fn is_gene_axis(&self, old: usize, new: usize) -> bool {
self.gene_axis
.as_ref()
.is_some_and(|g| old == g.n_old && new == g.new_to_old.len())
}
#[must_use]
pub fn reorders(&self, saved_dims: &[usize], fresh_dims: &[usize]) -> bool {
saved_dims
.iter()
.zip(fresh_dims)
.any(|(&o, &n)| self.is_gene_axis(o, n))
}
}
pub fn new_slab_value(name: &str, axis: Axis) -> Option<f64> {
if name.ends_with("topic.embeddings") {
return (axis == Axis::Embedding).then_some(0.0);
}
if name.ends_with("attn.query") || name.ends_with("fc.relu_linear_stack.0.weight") {
return match axis {
Axis::Embedding => Some(0.0),
Axis::Features => Some(0.0),
Axis::Topics => None,
};
}
if name.ends_with(crate::candle::feature_embedding::LOGITS_VAR_NAME) {
return (axis == Axis::Features).then_some(0.0);
}
if name.contains("z.mean") || name.contains("z.lnvar") {
return Some(if name.ends_with(".bias") {
NEW_TOPIC_LOGIT_BIAS
} else {
0.0
});
}
None
}
pub fn load_grown(parameters: &VarMap, path: &str, dims: &GrowthDims<'_>) -> anyhow::Result<()> {
let saved = candle_core::safetensors::load(path, &Device::Cpu)?;
let data = parameters.data().lock().expect("VarMap lock");
let (mut n_copied, mut n_grown) = (0usize, 0usize);
for (name, var) in data.iter() {
let s = saved.get(name).ok_or_else(|| {
anyhow::anyhow!("warm-start: {path} has no tensor named `{name}`; architectures differ")
})?;
let s = s.to_device(var.device())?;
let fresh = var.as_tensor();
if s.dims() == fresh.dims() && !dims.reorders(s.dims(), fresh.dims()) {
var.set(&s)?;
n_copied += 1;
continue;
}
var.set(&grow_tensor(name, fresh, &s, dims)?)?;
n_grown += 1;
}
log::info!("Warm-start: {n_copied} variables copied, {n_grown} grown");
Ok(())
}
pub fn grow_tensor(
name: &str,
fresh: &Tensor,
saved: &Tensor,
dims: &GrowthDims<'_>,
) -> anyhow::Result<Tensor> {
anyhow::ensure!(
saved.rank() == fresh.rank(),
"warm-start: `{name}` has rank {} in the checkpoint and {} here",
saved.rank(),
fresh.rank(),
);
let mut saved = saved.clone();
if let Some(g) = dims.gene_axis.as_ref() {
for dim in 0..saved.rank() {
if dims.is_gene_axis(saved.dims()[dim], fresh.dims()[dim]) {
saved = gather_gene_axis(name, &saved, dim, g)?;
}
}
}
let saved = &saved;
for (dim, (&old, &new)) in saved.dims().iter().zip(fresh.dims()).enumerate() {
anyhow::ensure!(
new >= old,
"warm-start: `{name}` shrank on axis {dim} ({old} → {new}); growth only adds capacity",
);
anyhow::ensure!(
old == new || dims.classify(old, new).is_some(),
"warm-start: `{name}` axis {dim} changed {old} → {new}, which is neither the \
requested K growth ({} → {}) nor H growth ({} → {}). That is an architecture \
change, not added capacity.",
dims.k_old,
dims.k_new,
dims.h_old,
dims.h_new,
);
}
let corner: Vec<std::ops::Range<usize>> = saved.dims().iter().map(|&d| 0..d).collect();
let mut out = fresh.slice_assign(&corner, saved)?;
for (dim, (&old, &new)) in saved.dims().iter().zip(fresh.dims()).enumerate() {
if old == new {
continue;
}
let Some(axis) = dims.classify(old, new) else {
continue;
};
if let Some(v) = new_slab_value(name, axis) {
let mut slab_shape: Vec<usize> = fresh.dims().to_vec();
slab_shape[dim] = new - old;
let slab = Tensor::full(v as f32, slab_shape.as_slice(), fresh.device())?
.to_dtype(fresh.dtype())?;
let mut r: Vec<std::ops::Range<usize>> = fresh.dims().iter().map(|&d| 0..d).collect();
r[dim] = old..new;
out = out.slice_assign(&r, &slab)?;
log::debug!(
"warm-start: `{name}` axis {dim} ({}) {old} → {new}, new slab = {v}",
axis.label(),
);
} else {
log::debug!(
"warm-start: `{name}` axis {dim} ({}) {old} → {new}, new slab keeps its init",
axis.label(),
);
}
}
Ok(out)
}
fn gather_gene_axis(
name: &str,
saved: &Tensor,
dim: usize,
remap: &AxisRemap<'_>,
) -> anyhow::Result<Tensor> {
let n_new = remap.new_to_old.len();
let n_known = remap.new_to_old.iter().filter(|p| p.is_some()).count();
let dev = saved.device();
log::debug!(
"warm-start: `{name}` axis {dim} (D) gathered {} → {n_new} ({n_known} known)",
remap.n_old,
);
let source = if n_known == n_new {
saved.clone()
} else {
let fill = match new_slab_value(name, Axis::Features) {
Some(v) => {
let mut one: Vec<usize> = saved.dims().to_vec();
one[dim] = 1;
Tensor::full(v as f32, one.as_slice(), dev)?.to_dtype(saved.dtype())?
}
None => saved.mean_keepdim(dim)?,
};
Tensor::cat(&[saved, &fill], dim)?
};
let unseen = remap.n_old as u32;
let idx: Vec<u32> = remap
.new_to_old
.iter()
.map(|p| p.map_or(unseen, |g| g as u32))
.collect();
Ok(source.index_select(&Tensor::from_vec(idx, n_new, dev)?, dim)?)
}