use anyhow::{bail, Context, Result};
use candle_core::{DType, Device, Module, Tensor};
use candle_nn::{
batch_norm, linear, BatchNorm, BatchNormConfig, Linear, ModuleT, VarBuilder, VarMap,
};
use finetype_model::char_cnn::{HierarchicalHead, HierarchyMap};
use finetype_model::model2vec_shared::Model2VecResources;
use finetype_model::value_attention::{ValueAttentionConfig, ValueAttentionPool};
use serde::{Deserialize, Serialize};
use std::io::{Read, Write};
use std::path::Path;
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub enum Activation {
#[default]
ReLU,
GELU,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub enum HeadType {
#[default]
Flat,
Hierarchical,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MultiBranchConfig {
pub char_dim: usize,
pub embed_dim: usize,
pub stats_dim: usize,
#[serde(default)]
pub header_dim: usize,
pub char_hidden: [usize; 2],
pub embed_hidden: [usize; 2],
pub stats_hidden: [usize; 2],
#[serde(default)]
pub header_hidden: [usize; 2],
#[serde(default)]
pub valid_dim: usize,
#[serde(default)]
pub valid_hidden: [usize; 2],
pub merge_hidden: [usize; 2],
pub n_classes: usize,
pub dropout: f32,
pub head_type: HeadType,
#[serde(default)]
pub activation: Activation,
#[serde(default)]
pub use_layer_norm: bool,
#[serde(default)]
pub value_attention: Option<ValueAttentionConfig>,
}
impl Default for MultiBranchConfig {
fn default() -> Self {
Self {
char_dim: 960,
embed_dim: 512,
stats_dim: 27,
header_dim: 128,
char_hidden: [300, 300],
embed_hidden: [200, 200],
stats_hidden: [128, 64],
header_hidden: [128, 64],
valid_dim: 0,
valid_hidden: [0, 0],
merge_hidden: [500, 500],
n_classes: 250,
dropout: 0.35,
head_type: HeadType::Flat,
activation: Activation::ReLU,
use_layer_norm: false,
value_attention: None,
}
}
}
impl MultiBranchConfig {
pub fn has_header_branch(&self) -> bool {
self.header_dim > 0 && self.header_hidden[0] > 0 && self.header_hidden[1] > 0
}
pub fn has_validation_branch(&self) -> bool {
self.valid_dim > 0 && self.valid_hidden[0] > 0 && self.valid_hidden[1] > 0
}
pub fn embed_branch_input_dim(&self) -> usize {
match &self.value_attention {
Some(va) => {
let blender = if va.keep_blender_concat {
self.embed_dim
} else {
0
};
blender + va.output_dim()
}
None => self.embed_dim,
}
}
pub fn merged_dim(&self) -> usize {
let mut dim = self.char_hidden[1] + self.embed_hidden[1] + self.stats_hidden[1];
if self.has_header_branch() {
dim += self.header_hidden[1];
}
if self.has_validation_branch() {
dim += self.valid_hidden[1];
}
dim
}
pub fn save(&self, path: &Path) -> Result<()> {
let json = serde_json::to_string_pretty(self)?;
std::fs::write(path, json)?;
Ok(())
}
pub fn load(path: &Path) -> Result<Self> {
let json = std::fs::read_to_string(path)?;
let config: Self = serde_json::from_str(&json)?;
Ok(config)
}
}
struct BranchWeights {
input_norm: Option<candle_nn::LayerNorm>,
linear1: Linear,
linear2: Linear,
dropout: f32,
activation: Activation,
}
impl BranchWeights {
fn new(
input_dim: usize,
hidden: [usize; 2],
dropout: f32,
activation: &Activation,
vb: VarBuilder,
) -> candle_core::Result<Self> {
Self::new_inner(input_dim, hidden, dropout, false, activation, vb)
}
fn new_with_input_norm(
input_dim: usize,
hidden: [usize; 2],
dropout: f32,
activation: &Activation,
vb: VarBuilder,
) -> candle_core::Result<Self> {
Self::new_inner(input_dim, hidden, dropout, true, activation, vb)
}
fn new_inner(
input_dim: usize,
hidden: [usize; 2],
dropout: f32,
normalize_input: bool,
activation: &Activation,
vb: VarBuilder,
) -> candle_core::Result<Self> {
let input_norm = if normalize_input {
Some(candle_nn::layer_norm(
input_dim,
candle_nn::LayerNormConfig::default(),
vb.pp("input_ln"),
)?)
} else {
None
};
let linear1 = linear(input_dim, hidden[0], vb.pp("l1"))?;
let linear2 = linear(hidden[0], hidden[1], vb.pp("l2"))?;
Ok(Self {
input_norm,
linear1,
linear2,
dropout,
activation: activation.clone(),
})
}
fn activate(&self, x: &Tensor) -> candle_core::Result<Tensor> {
match self.activation {
Activation::ReLU => x.relu(),
Activation::GELU => x.gelu_erf(),
}
}
fn forward(&self, x: &Tensor, train: bool) -> candle_core::Result<Tensor> {
let x = match &self.input_norm {
Some(ln) => ln.forward(x)?,
None => x.clone(),
};
let h = self.linear1.forward_t(&x, false)?;
let h = self.activate(&h)?;
let h = if train {
crate::seeded_rng::dropout(&h, self.dropout)?
} else {
h
};
let h = self.linear2.forward_t(&h, false)?;
let h = self.activate(&h)?;
if train {
crate::seeded_rng::dropout(&h, self.dropout)
} else {
Ok(h)
}
}
}
enum MergeNorm {
Batch(BatchNorm),
Layer(candle_nn::LayerNorm),
}
pub struct MultiBranchModel {
char_branch: BranchWeights,
embed_branch: BranchWeights,
stats_branch: BranchWeights,
header_branch: Option<BranchWeights>,
valid_branch: Option<BranchWeights>,
value_attention: Option<ValueAttentionPool>,
merge_norm: MergeNorm,
merge_linear1: Linear,
merge_linear2: Linear,
head: Option<Linear>,
hierarchical: Option<HierarchicalHead>,
config: MultiBranchConfig,
}
impl MultiBranchModel {
pub fn new(config: &MultiBranchConfig, vb: VarBuilder) -> candle_core::Result<Self> {
let (
char_branch,
embed_branch,
stats_branch,
header_branch,
valid_branch,
value_attention,
merge_norm,
merge_linear1,
merge_linear2,
) = Self::build_trunk(config, &vb)?;
let head = linear(config.merge_hidden[1], config.n_classes, vb.pp("head"))?;
Ok(Self {
char_branch,
embed_branch,
stats_branch,
header_branch,
valid_branch,
value_attention,
merge_norm,
merge_linear1,
merge_linear2,
head: Some(head),
hierarchical: None,
config: config.clone(),
})
}
pub fn new_hierarchical(
config: &MultiBranchConfig,
labels: &[String],
vb: VarBuilder,
) -> candle_core::Result<Self> {
let (
char_branch,
embed_branch,
stats_branch,
header_branch,
valid_branch,
value_attention,
merge_norm,
merge_linear1,
merge_linear2,
) = Self::build_trunk(config, &vb)?;
let hier_head = HierarchicalHead::new(
config.merge_hidden[1],
labels,
vb.pp(HierarchicalHead::VARBUILDER_PREFIX),
)?;
Ok(Self {
char_branch,
embed_branch,
stats_branch,
header_branch,
valid_branch,
value_attention,
merge_norm,
merge_linear1,
merge_linear2,
head: None,
hierarchical: Some(hier_head),
config: config.clone(),
})
}
#[allow(clippy::type_complexity)]
fn build_trunk(
config: &MultiBranchConfig,
vb: &VarBuilder,
) -> candle_core::Result<(
BranchWeights,
BranchWeights,
BranchWeights,
Option<BranchWeights>,
Option<BranchWeights>,
Option<ValueAttentionPool>,
MergeNorm,
Linear,
Linear,
)> {
let build_branch = |input_dim, hidden, name: &str| -> candle_core::Result<BranchWeights> {
if config.use_layer_norm {
BranchWeights::new_with_input_norm(
input_dim,
hidden,
config.dropout,
&config.activation,
vb.pp(name),
)
} else {
BranchWeights::new(
input_dim,
hidden,
config.dropout,
&config.activation,
vb.pp(name),
)
}
};
let char_branch = build_branch(config.char_dim, config.char_hidden, "char")?;
let embed_branch = build_branch(
config.embed_branch_input_dim(),
config.embed_hidden,
"embed",
)?;
let stats_branch = build_branch(config.stats_dim, config.stats_hidden, "stats")?;
let value_attention = match &config.value_attention {
Some(va) => Some(ValueAttentionPool::new(va, vb.pp("value_attn"))?),
None => None,
};
let header_branch = if config.has_header_branch() {
Some(BranchWeights::new_with_input_norm(
config.header_dim,
config.header_hidden,
config.dropout,
&config.activation,
vb.pp("header"),
)?)
} else {
None
};
let valid_branch = if config.has_validation_branch() {
Some(BranchWeights::new(
config.valid_dim,
config.valid_hidden,
config.dropout,
&config.activation,
vb.pp("valid"),
)?)
} else {
None
};
let merged_dim = config.merged_dim();
let merge_norm = if config.use_layer_norm {
MergeNorm::Layer(candle_nn::layer_norm(
merged_dim,
candle_nn::LayerNormConfig::default(),
vb.pp("merge_ln"),
)?)
} else {
MergeNorm::Batch(batch_norm(
merged_dim,
BatchNormConfig::default(),
vb.pp("merge_bn"),
)?)
};
let merge_linear1 = linear(merged_dim, config.merge_hidden[0], vb.pp("merge_l1"))?;
let merge_linear2 = linear(
config.merge_hidden[0],
config.merge_hidden[1],
vb.pp("merge_l2"),
)?;
Ok((
char_branch,
embed_branch,
stats_branch,
header_branch,
valid_branch,
value_attention,
merge_norm,
merge_linear1,
merge_linear2,
))
}
pub fn has_value_attention(&self) -> bool {
self.value_attention.is_some()
}
pub fn embed_input(
&self,
blender: &Tensor,
value_embeds: Option<&Tensor>,
value_mask: Option<&Tensor>,
train: bool,
) -> candle_core::Result<Tensor> {
let Some(pool) = &self.value_attention else {
return Ok(blender.clone());
};
let (embeds, mask) = value_embeds.zip(value_mask).ok_or_else(|| {
candle_core::Error::Msg(
"value-attention model requires value_embeds + value_mask".into(),
)
})?;
let pooled = pool.forward(embeds, mask, train)?;
let keep = self
.config
.value_attention
.as_ref()
.map(|v| v.keep_blender_concat)
.unwrap_or(false);
if keep {
Tensor::cat(&[blender, &pooled], 1)
} else {
Ok(pooled)
}
}
pub fn forward_trunk(
&self,
char_feats: &Tensor,
embed_feats: &Tensor,
stats_feats: &Tensor,
header_feats: Option<&Tensor>,
valid_feats: Option<&Tensor>,
train: bool,
) -> candle_core::Result<Tensor> {
let char_out = self.char_branch.forward(char_feats, train)?;
let embed_out = self.embed_branch.forward(embed_feats, train)?;
let stats_out = self.stats_branch.forward(stats_feats, train)?;
let mut branch_outputs = vec![char_out, embed_out, stats_out];
let batch_size = branch_outputs[0].dim(0)?;
if let Some(ref hb) = &self.header_branch {
let header_input = match header_feats {
Some(hf) => hf.clone(),
None => Tensor::zeros(
(batch_size, self.config.header_dim),
DType::F32,
char_feats.device(),
)?,
};
let header_out = hb.forward(&header_input, train)?;
branch_outputs.push(header_out);
}
if let Some(ref vb_weights) = &self.valid_branch {
let valid_input = match valid_feats {
Some(vf) => vf.clone(),
None => Tensor::zeros(
(batch_size, self.config.valid_dim),
DType::F32,
char_feats.device(),
)?,
};
let valid_out = vb_weights.forward(&valid_input, train)?;
branch_outputs.push(valid_out);
}
let merged = Tensor::cat(&branch_outputs, 1)?;
let normed = match &self.merge_norm {
MergeNorm::Batch(bn) => {
let merged_3d = merged.unsqueeze(2)?; let normed_3d = bn.forward_t(&merged_3d, train)?;
normed_3d.squeeze(2)? }
MergeNorm::Layer(ln) => {
ln.forward(&merged)?
}
};
let activate = |x: &Tensor| -> candle_core::Result<Tensor> {
match self.config.activation {
Activation::ReLU => x.relu(),
Activation::GELU => x.gelu_erf(),
}
};
let h = self.merge_linear1.forward_t(&normed, false)?;
let h = activate(&h)?;
let h = if train {
crate::seeded_rng::dropout(&h, self.config.dropout)?
} else {
h
};
let h = self.merge_linear2.forward_t(&h, false)?;
let h = activate(&h)?;
if train {
crate::seeded_rng::dropout(&h, self.config.dropout)
} else {
Ok(h)
}
}
pub fn forward(
&self,
char_feats: &Tensor,
embed_feats: &Tensor,
stats_feats: &Tensor,
header_feats: Option<&Tensor>,
valid_feats: Option<&Tensor>,
train: bool,
) -> candle_core::Result<Tensor> {
let hidden = self.forward_trunk(
char_feats,
embed_feats,
stats_feats,
header_feats,
valid_feats,
train,
)?;
if let Some(ref head) = self.head {
head.forward_t(&hidden, false)
} else if let Some(ref hier) = self.hierarchical {
hier.forward(&hidden, self.config.n_classes)
} else {
candle_core::bail!("No classification head configured");
}
}
#[allow(clippy::type_complexity)]
pub fn forward_levels(
&self,
char_feats: &Tensor,
embed_feats: &Tensor,
stats_feats: &Tensor,
header_feats: Option<&Tensor>,
valid_feats: Option<&Tensor>,
train: bool,
) -> candle_core::Result<(Tensor, Vec<Tensor>, Vec<Vec<Option<Tensor>>>)> {
let hidden = self.forward_trunk(
char_feats,
embed_feats,
stats_feats,
header_feats,
valid_feats,
train,
)?;
match &self.hierarchical {
Some(ref hier) => hier.forward_levels(&hidden),
None => candle_core::bail!("forward_levels() requires a hierarchical head"),
}
}
pub fn hierarchical_head(&self) -> Option<&HierarchicalHead> {
self.hierarchical.as_ref()
}
pub fn is_hierarchical(&self) -> bool {
self.hierarchical.is_some()
}
pub fn config(&self) -> &MultiBranchConfig {
&self.config
}
}
const FTMB_MAGIC: &[u8; 4] = b"FTMB";
const FTMB_VERSION_V2: u32 = 2;
const FTMB_VERSION_V3: u32 = 3;
const FTMB_HEADER_SIZE_V2: usize = 24;
#[allow(dead_code)]
const FTMB_HEADER_SIZE_V3: usize = 28;
const FTMB_VERSION_V4: u32 = 4;
pub const FTMB_HEADER_SIZE_V4: usize = 30;
const FTMB_VERSION_V5: u32 = 5;
const FTMB_VERSION_V6: u32 = 6;
#[derive(Debug, Clone, Default)]
pub struct TrainingRecord {
pub label: String,
pub char_features: Vec<f32>,
pub embed_features: Vec<f32>,
pub stats_features: Vec<f32>,
pub header_features: Vec<f32>,
pub validation_features: Vec<f32>,
pub values: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct TableGroup {
pub record_indices: Vec<usize>,
pub sibling_headers: Vec<String>,
}
pub fn write_training_data(
path: &Path,
records: &[TrainingRecord],
char_dim: u16,
embed_dim: u16,
stats_dim: u16,
header_dim: u16,
) -> Result<()> {
let mut file = std::fs::File::create(path)
.with_context(|| format!("Failed to create training data file: {}", path.display()))?;
file.write_all(FTMB_MAGIC)?;
file.write_all(&FTMB_VERSION_V2.to_le_bytes())?;
file.write_all(&(records.len() as u64).to_le_bytes())?;
file.write_all(&char_dim.to_le_bytes())?;
file.write_all(&embed_dim.to_le_bytes())?;
file.write_all(&stats_dim.to_le_bytes())?;
file.write_all(&header_dim.to_le_bytes())?;
for record in records {
let label_bytes = record.label.as_bytes();
if label_bytes.len() > u16::MAX as usize {
bail!(
"Label too long ({} bytes): {}",
label_bytes.len(),
record.label
);
}
file.write_all(&(label_bytes.len() as u16).to_le_bytes())?;
file.write_all(label_bytes)?;
if record.char_features.len() != char_dim as usize {
bail!(
"char_features length {} != expected {}",
record.char_features.len(),
char_dim
);
}
if record.embed_features.len() != embed_dim as usize {
bail!(
"embed_features length {} != expected {}",
record.embed_features.len(),
embed_dim
);
}
if record.stats_features.len() != stats_dim as usize {
bail!(
"stats_features length {} != expected {}",
record.stats_features.len(),
stats_dim
);
}
if record.header_features.len() != header_dim as usize {
bail!(
"header_features length {} != expected {}",
record.header_features.len(),
header_dim
);
}
for &v in &record.char_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.embed_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.stats_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.header_features {
file.write_all(&v.to_le_bytes())?;
}
}
Ok(())
}
#[derive(Debug)]
pub struct FtmbHeader {
pub version: u32,
pub n_records: u64,
pub char_dim: u16,
pub embed_dim: u16,
pub stats_dim: u16,
pub header_dim: u16,
pub n_groups: u16,
pub valid_dim: u16,
}
pub fn read_training_header(path: &Path) -> Result<FtmbHeader> {
let mut file = std::fs::File::open(path)
.with_context(|| format!("Failed to open training data file: {}", path.display()))?;
let mut header = [0u8; FTMB_HEADER_SIZE_V4];
file.read_exact(&mut header[..FTMB_HEADER_SIZE_V2])
.context("Failed to read FTMB header")?;
if &header[0..4] != FTMB_MAGIC {
bail!(
"Invalid FTMB magic: expected {:?}, got {:?}",
FTMB_MAGIC,
&header[0..4]
);
}
let version = u32::from_le_bytes(header[4..8].try_into().unwrap());
if version != 1
&& version != 2
&& version != 3
&& version != 4
&& version != FTMB_VERSION_V5
&& version != FTMB_VERSION_V6
{
bail!(
"Unsupported FTMB version: {} (expected 1, 2, 3, 4, 5, or 6)",
version
);
}
let n_records = u64::from_le_bytes(header[8..16].try_into().unwrap());
let char_dim = u16::from_le_bytes(header[16..18].try_into().unwrap());
let embed_dim = u16::from_le_bytes(header[18..20].try_into().unwrap());
let stats_dim = u16::from_le_bytes(header[20..22].try_into().unwrap());
let header_dim = if version >= 2 {
u16::from_le_bytes(header[22..24].try_into().unwrap())
} else {
0
};
let n_groups = if version >= 3 {
file.read_exact(&mut header[24..28])
.context("Failed to read v3 header extension")?;
u16::from_le_bytes(header[24..26].try_into().unwrap())
} else {
0
};
let valid_dim = if version >= 4 {
file.read_exact(&mut header[28..30])
.context("Failed to read v4 header extension")?;
u16::from_le_bytes(header[28..30].try_into().unwrap())
} else {
0
};
Ok(FtmbHeader {
version,
n_records,
char_dim,
embed_dim,
stats_dim,
header_dim,
n_groups,
valid_dim,
})
}
pub fn read_training_data(
path: &Path,
) -> Result<(FtmbHeader, Vec<TrainingRecord>, Vec<TableGroup>)> {
let mut file = std::fs::File::open(path)
.with_context(|| format!("Failed to open training data file: {}", path.display()))?;
let mut header_buf = [0u8; FTMB_HEADER_SIZE_V4];
file.read_exact(&mut header_buf[..FTMB_HEADER_SIZE_V2])
.context("Failed to read FTMB header")?;
if &header_buf[0..4] != FTMB_MAGIC {
bail!("Invalid FTMB magic");
}
let version = u32::from_le_bytes(header_buf[4..8].try_into().unwrap());
if version != 1
&& version != 2
&& version != 3
&& version != 4
&& version != FTMB_VERSION_V5
&& version != FTMB_VERSION_V6
{
bail!(
"Unsupported FTMB version: {} (expected 1, 2, 3, 4, 5, or 6)",
version
);
}
let n_records = u64::from_le_bytes(header_buf[8..16].try_into().unwrap());
let char_dim = u16::from_le_bytes(header_buf[16..18].try_into().unwrap()) as usize;
let embed_dim = u16::from_le_bytes(header_buf[18..20].try_into().unwrap()) as usize;
let stats_dim = u16::from_le_bytes(header_buf[20..22].try_into().unwrap()) as usize;
let header_dim = if version >= 2 {
u16::from_le_bytes(header_buf[22..24].try_into().unwrap()) as usize
} else {
0
};
let n_groups = if version >= 3 {
file.read_exact(&mut header_buf[24..28])
.context("Failed to read v3 header extension")?;
u16::from_le_bytes(header_buf[24..26].try_into().unwrap()) as usize
} else {
0
};
let valid_dim = if version >= 4 {
file.read_exact(&mut header_buf[28..30])
.context("Failed to read v4 header extension")?;
u16::from_le_bytes(header_buf[28..30].try_into().unwrap()) as usize
} else {
0
};
let header = FtmbHeader {
version,
n_records,
char_dim: char_dim as u16,
embed_dim: embed_dim as u16,
stats_dim: stats_dim as u16,
header_dim: header_dim as u16,
n_groups: n_groups as u16,
valid_dim: valid_dim as u16,
};
let mut records = Vec::with_capacity(n_records as usize);
let mut table_groups = Vec::new();
let mut label_len_buf = [0u8; 2];
let mut f32_buf = [0u8; 4];
if version >= 3 {
let mut record_offset = 0usize;
for _ in 0..n_groups {
let mut group_header = [0u8; 4];
file.read_exact(&mut group_header)?;
let n_columns = u16::from_le_bytes(group_header[0..2].try_into().unwrap()) as usize;
let n_sibling_headers =
u16::from_le_bytes(group_header[2..4].try_into().unwrap()) as usize;
let mut sibling_headers = Vec::with_capacity(n_sibling_headers);
for _ in 0..n_sibling_headers {
file.read_exact(&mut label_len_buf)?;
let header_len = u16::from_le_bytes(label_len_buf) as usize;
let mut header_bytes = vec![0u8; header_len];
file.read_exact(&mut header_bytes)?;
sibling_headers
.push(String::from_utf8(header_bytes).context("Invalid UTF-8 in header")?);
}
let mut group_indices = Vec::with_capacity(n_columns);
for _ in 0..n_columns {
file.read_exact(&mut label_len_buf)?;
let label_len = u16::from_le_bytes(label_len_buf) as usize;
let mut label_buf = vec![0u8; label_len];
file.read_exact(&mut label_buf)?;
let label = String::from_utf8(label_buf).context("Invalid UTF-8 in label")?;
let mut col_idx_buf = [0u8; 2];
file.read_exact(&mut col_idx_buf)?;
let _column_index = u16::from_le_bytes(col_idx_buf);
let mut char_features = Vec::with_capacity(char_dim);
for _ in 0..char_dim {
file.read_exact(&mut f32_buf)?;
char_features.push(f32::from_le_bytes(f32_buf));
}
let mut embed_features = Vec::with_capacity(embed_dim);
for _ in 0..embed_dim {
file.read_exact(&mut f32_buf)?;
embed_features.push(f32::from_le_bytes(f32_buf));
}
let mut stats_features = Vec::with_capacity(stats_dim);
for _ in 0..stats_dim {
file.read_exact(&mut f32_buf)?;
stats_features.push(f32::from_le_bytes(f32_buf));
}
let mut header_features = Vec::with_capacity(header_dim);
for _ in 0..header_dim {
file.read_exact(&mut f32_buf)?;
header_features.push(f32::from_le_bytes(f32_buf));
}
let validation_features = if version >= 4 && valid_dim > 0 {
let mut vf = Vec::with_capacity(valid_dim);
for _ in 0..valid_dim {
file.read_exact(&mut f32_buf)?;
vf.push(f32::from_le_bytes(f32_buf));
}
vf
} else {
Vec::new()
};
let values = if version >= FTMB_VERSION_V6 {
let mut n_values_buf = [0u8; 2];
file.read_exact(&mut n_values_buf)?;
let n_values = u16::from_le_bytes(n_values_buf) as usize;
let mut vals = Vec::with_capacity(n_values);
for _ in 0..n_values {
file.read_exact(&mut label_len_buf)?;
let val_len = u16::from_le_bytes(label_len_buf) as usize;
let mut val_buf = vec![0u8; val_len];
file.read_exact(&mut val_buf)?;
vals.push(String::from_utf8(val_buf).context("Invalid UTF-8 in value")?);
}
vals
} else {
Vec::new()
};
group_indices.push(record_offset);
records.push(TrainingRecord {
label,
char_features,
embed_features,
stats_features,
header_features,
validation_features,
values,
});
record_offset += 1;
}
table_groups.push(TableGroup {
record_indices: group_indices,
sibling_headers,
});
}
} else {
for _ in 0..n_records {
file.read_exact(&mut label_len_buf)?;
let label_len = u16::from_le_bytes(label_len_buf) as usize;
let mut label_buf = vec![0u8; label_len];
file.read_exact(&mut label_buf)?;
let label = String::from_utf8(label_buf).context("Invalid UTF-8 in label")?;
let mut char_features = Vec::with_capacity(char_dim);
for _ in 0..char_dim {
file.read_exact(&mut f32_buf)?;
char_features.push(f32::from_le_bytes(f32_buf));
}
let mut embed_features = Vec::with_capacity(embed_dim);
for _ in 0..embed_dim {
file.read_exact(&mut f32_buf)?;
embed_features.push(f32::from_le_bytes(f32_buf));
}
let mut stats_features = Vec::with_capacity(stats_dim);
for _ in 0..stats_dim {
file.read_exact(&mut f32_buf)?;
stats_features.push(f32::from_le_bytes(f32_buf));
}
let header_features = if version >= 2 && header_dim > 0 {
let mut hf = Vec::with_capacity(header_dim);
for _ in 0..header_dim {
file.read_exact(&mut f32_buf)?;
hf.push(f32::from_le_bytes(f32_buf));
}
hf
} else {
vec![0.0f32; header_dim]
};
let validation_features = if version >= 4 && valid_dim > 0 {
let mut vf = Vec::with_capacity(valid_dim);
for _ in 0..valid_dim {
file.read_exact(&mut f32_buf)?;
vf.push(f32::from_le_bytes(f32_buf));
}
vf
} else {
Vec::new()
};
records.push(TrainingRecord {
label,
char_features,
embed_features,
stats_features,
header_features,
validation_features,
values: Vec::new(),
});
}
if !records.is_empty() {
table_groups.push(TableGroup {
record_indices: (0..records.len()).collect(),
sibling_headers: Vec::new(),
});
}
}
Ok((header, records, table_groups))
}
pub fn write_training_data_v3(
path: &Path,
records: &[TrainingRecord],
table_groups: &[TableGroup],
char_dim: u16,
embed_dim: u16,
stats_dim: u16,
header_dim: u16,
) -> Result<()> {
let mut file = std::fs::File::create(path)
.with_context(|| format!("Failed to create training data file: {}", path.display()))?;
let total_records: usize = table_groups.iter().map(|g| g.record_indices.len()).sum();
if total_records != records.len() {
bail!(
"Table groups reference {} records but {} were provided",
total_records,
records.len()
);
}
file.write_all(FTMB_MAGIC)?;
file.write_all(&FTMB_VERSION_V3.to_le_bytes())?;
file.write_all(&(records.len() as u64).to_le_bytes())?;
file.write_all(&char_dim.to_le_bytes())?;
file.write_all(&embed_dim.to_le_bytes())?;
file.write_all(&stats_dim.to_le_bytes())?;
file.write_all(&header_dim.to_le_bytes())?;
file.write_all(&(table_groups.len() as u16).to_le_bytes())?;
file.write_all(&[0u8; 2])?;
for group in table_groups {
file.write_all(&(group.record_indices.len() as u16).to_le_bytes())?;
file.write_all(&(group.sibling_headers.len() as u16).to_le_bytes())?;
for header_name in &group.sibling_headers {
let header_bytes = header_name.as_bytes();
if header_bytes.len() > u16::MAX as usize {
bail!(
"Sibling header too long ({} bytes): {}",
header_bytes.len(),
header_name
);
}
file.write_all(&(header_bytes.len() as u16).to_le_bytes())?;
file.write_all(header_bytes)?;
}
for (col_idx, &record_idx) in group.record_indices.iter().enumerate() {
let record = &records[record_idx];
let label_bytes = record.label.as_bytes();
if label_bytes.len() > u16::MAX as usize {
bail!(
"Label too long ({} bytes): {}",
label_bytes.len(),
record.label
);
}
file.write_all(&(label_bytes.len() as u16).to_le_bytes())?;
file.write_all(label_bytes)?;
file.write_all(&(col_idx as u16).to_le_bytes())?;
if record.char_features.len() != char_dim as usize {
bail!(
"char_features length {} != expected {}",
record.char_features.len(),
char_dim
);
}
if record.embed_features.len() != embed_dim as usize {
bail!(
"embed_features length {} != expected {}",
record.embed_features.len(),
embed_dim
);
}
if record.stats_features.len() != stats_dim as usize {
bail!(
"stats_features length {} != expected {}",
record.stats_features.len(),
stats_dim
);
}
if record.header_features.len() != header_dim as usize {
bail!(
"header_features length {} != expected {}",
record.header_features.len(),
header_dim
);
}
for &v in &record.char_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.embed_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.stats_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.header_features {
file.write_all(&v.to_le_bytes())?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn write_training_data_v4(
path: &Path,
records: &[TrainingRecord],
table_groups: &[TableGroup],
char_dim: u16,
embed_dim: u16,
stats_dim: u16,
header_dim: u16,
valid_dim: u16,
) -> Result<()> {
let mut file = std::fs::File::create(path)
.with_context(|| format!("Failed to create training data file: {}", path.display()))?;
let total_records: usize = table_groups.iter().map(|g| g.record_indices.len()).sum();
if total_records != records.len() {
bail!(
"Table groups reference {} records but {} were provided",
total_records,
records.len()
);
}
file.write_all(FTMB_MAGIC)?;
file.write_all(&FTMB_VERSION_V4.to_le_bytes())?;
file.write_all(&(records.len() as u64).to_le_bytes())?;
file.write_all(&char_dim.to_le_bytes())?;
file.write_all(&embed_dim.to_le_bytes())?;
file.write_all(&stats_dim.to_le_bytes())?;
file.write_all(&header_dim.to_le_bytes())?;
file.write_all(&(table_groups.len() as u16).to_le_bytes())?;
file.write_all(&[0u8; 2])?; file.write_all(&valid_dim.to_le_bytes())?;
for group in table_groups {
file.write_all(&(group.record_indices.len() as u16).to_le_bytes())?;
file.write_all(&(group.sibling_headers.len() as u16).to_le_bytes())?;
for header_name in &group.sibling_headers {
let header_bytes = header_name.as_bytes();
if header_bytes.len() > u16::MAX as usize {
bail!(
"Sibling header too long ({} bytes): {}",
header_bytes.len(),
header_name
);
}
file.write_all(&(header_bytes.len() as u16).to_le_bytes())?;
file.write_all(header_bytes)?;
}
for (col_idx, &record_idx) in group.record_indices.iter().enumerate() {
let record = &records[record_idx];
let label_bytes = record.label.as_bytes();
if label_bytes.len() > u16::MAX as usize {
bail!(
"Label too long ({} bytes): {}",
label_bytes.len(),
record.label
);
}
file.write_all(&(label_bytes.len() as u16).to_le_bytes())?;
file.write_all(label_bytes)?;
file.write_all(&(col_idx as u16).to_le_bytes())?;
if record.char_features.len() != char_dim as usize {
bail!(
"char_features length {} != expected {}",
record.char_features.len(),
char_dim
);
}
if record.embed_features.len() != embed_dim as usize {
bail!(
"embed_features length {} != expected {}",
record.embed_features.len(),
embed_dim
);
}
if record.stats_features.len() != stats_dim as usize {
bail!(
"stats_features length {} != expected {}",
record.stats_features.len(),
stats_dim
);
}
if record.header_features.len() != header_dim as usize {
bail!(
"header_features length {} != expected {}",
record.header_features.len(),
header_dim
);
}
if record.validation_features.len() != valid_dim as usize {
bail!(
"validation_features length {} != expected {}",
record.validation_features.len(),
valid_dim
);
}
for &v in &record.char_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.embed_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.stats_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.header_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.validation_features {
file.write_all(&v.to_le_bytes())?;
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn write_training_data_v6(
path: &Path,
records: &[TrainingRecord],
table_groups: &[TableGroup],
char_dim: u16,
embed_dim: u16,
stats_dim: u16,
header_dim: u16,
valid_dim: u16,
n_values: u16,
) -> Result<()> {
let mut file = std::fs::File::create(path)
.with_context(|| format!("Failed to create training data file: {}", path.display()))?;
let total_records: usize = table_groups.iter().map(|g| g.record_indices.len()).sum();
if total_records != records.len() {
bail!(
"Table groups reference {} records but {} were provided",
total_records,
records.len()
);
}
file.write_all(FTMB_MAGIC)?;
file.write_all(&FTMB_VERSION_V6.to_le_bytes())?;
file.write_all(&(records.len() as u64).to_le_bytes())?;
file.write_all(&char_dim.to_le_bytes())?;
file.write_all(&embed_dim.to_le_bytes())?;
file.write_all(&stats_dim.to_le_bytes())?;
file.write_all(&header_dim.to_le_bytes())?;
file.write_all(&(table_groups.len() as u16).to_le_bytes())?;
file.write_all(&[0u8; 2])?; file.write_all(&valid_dim.to_le_bytes())?;
for group in table_groups {
file.write_all(&(group.record_indices.len() as u16).to_le_bytes())?;
file.write_all(&(group.sibling_headers.len() as u16).to_le_bytes())?;
for header_name in &group.sibling_headers {
let header_bytes = header_name.as_bytes();
if header_bytes.len() > u16::MAX as usize {
bail!(
"Sibling header too long ({} bytes): {}",
header_bytes.len(),
header_name
);
}
file.write_all(&(header_bytes.len() as u16).to_le_bytes())?;
file.write_all(header_bytes)?;
}
for (col_idx, &record_idx) in group.record_indices.iter().enumerate() {
let record = &records[record_idx];
let label_bytes = record.label.as_bytes();
if label_bytes.len() > u16::MAX as usize {
bail!(
"Label too long ({} bytes): {}",
label_bytes.len(),
record.label
);
}
file.write_all(&(label_bytes.len() as u16).to_le_bytes())?;
file.write_all(label_bytes)?;
file.write_all(&(col_idx as u16).to_le_bytes())?;
for (name, got, want) in [
(
"char_features",
record.char_features.len(),
char_dim as usize,
),
(
"embed_features",
record.embed_features.len(),
embed_dim as usize,
),
(
"stats_features",
record.stats_features.len(),
stats_dim as usize,
),
(
"header_features",
record.header_features.len(),
header_dim as usize,
),
(
"validation_features",
record.validation_features.len(),
valid_dim as usize,
),
] {
if got != want {
bail!("{name} length {got} != expected {want}");
}
}
for &v in &record.char_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.embed_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.stats_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.header_features {
file.write_all(&v.to_le_bytes())?;
}
for &v in &record.validation_features {
file.write_all(&v.to_le_bytes())?;
}
let n = (record.values.len()).min(n_values as usize);
file.write_all(&(n as u16).to_le_bytes())?;
for val in record.values.iter().take(n) {
let val: &str = if val.len() > u16::MAX as usize {
let mut cut = u16::MAX as usize;
while cut > 0 && !val.is_char_boundary(cut) {
cut -= 1;
}
&val[..cut]
} else {
val
};
let vb = val.as_bytes();
file.write_all(&(vb.len() as u16).to_le_bytes())?;
file.write_all(vb)?;
}
}
}
Ok(())
}
pub struct FrozenSiblingContext {
blocks: Vec<FrozenTransformerBlock>,
final_norm_weight: Tensor,
final_norm_bias: Tensor,
embed_dim: usize,
}
struct FrozenTransformerBlock {
norm1_weight: Tensor,
norm1_bias: Tensor,
wq: Tensor,
bq: Tensor,
wk: Tensor,
bk: Tensor,
wv: Tensor,
bv: Tensor,
out_weight: Tensor,
out_bias: Tensor,
n_heads: usize,
head_dim: usize,
norm2_weight: Tensor,
norm2_bias: Tensor,
ffn_w1: Tensor,
ffn_b1: Tensor,
ffn_w2: Tensor,
ffn_b2: Tensor,
}
impl FrozenSiblingContext {
pub fn load(model_dir: &Path, device: &Device) -> Result<Self> {
use finetype_model::sibling_context::SiblingContextConfig;
let config_bytes =
std::fs::read(model_dir.join("config.json")).context("Failed to read config.json")?;
let config: SiblingContextConfig =
serde_json::from_slice(&config_bytes).context("Failed to parse config.json")?;
let tensors = candle_core::safetensors::load(model_dir.join("model.safetensors"), device)?;
let get = |name: &str| -> Result<Tensor> {
tensors
.get(name)
.cloned()
.ok_or_else(|| {
anyhow::anyhow!("Missing tensor '{}' in sibling-context model", name)
})
.and_then(|t| Ok(t.to_dtype(DType::F32)?))
};
let d = config.embed_dim;
let n_heads = config.n_heads;
let head_dim = d / n_heads;
let mut blocks = Vec::with_capacity(config.n_layers);
for i in 0..config.n_layers {
let p = format!("blocks.{}", i);
blocks.push(FrozenTransformerBlock {
norm1_weight: get(&format!("{p}.norm1.weight"))?,
norm1_bias: get(&format!("{p}.norm1.bias"))?,
wq: get(&format!("{p}.attn.wq"))?,
bq: get(&format!("{p}.attn.bq"))?,
wk: get(&format!("{p}.attn.wk"))?,
bk: get(&format!("{p}.attn.bk"))?,
wv: get(&format!("{p}.attn.wv"))?,
bv: get(&format!("{p}.attn.bv"))?,
out_weight: get(&format!("{p}.attn.out_weight"))?,
out_bias: get(&format!("{p}.attn.out_bias"))?,
n_heads,
head_dim,
norm2_weight: get(&format!("{p}.norm2.weight"))?,
norm2_bias: get(&format!("{p}.norm2.bias"))?,
ffn_w1: get(&format!("{p}.ffn.w1"))?,
ffn_b1: get(&format!("{p}.ffn.b1"))?,
ffn_w2: get(&format!("{p}.ffn.w2"))?,
ffn_b2: get(&format!("{p}.ffn.b2"))?,
});
}
Ok(Self {
blocks,
final_norm_weight: get("final_norm.weight")?,
final_norm_bias: get("final_norm.bias")?,
embed_dim: d,
})
}
pub fn forward(&self, header_embeds: &Tensor) -> Result<Tensor> {
let mut out = header_embeds.clone();
for block in &self.blocks {
out = block.forward(&out)?;
}
frozen_layer_norm(&out, &self.final_norm_weight, &self.final_norm_bias)
}
pub fn embed_dim(&self) -> usize {
self.embed_dim
}
}
impl FrozenTransformerBlock {
fn forward(&self, x: &Tensor) -> Result<Tensor> {
let normed = frozen_layer_norm(x, &self.norm1_weight, &self.norm1_bias)?;
let attn_out = self.forward_attn(&normed)?;
let x = (x + &attn_out)?;
let normed = frozen_layer_norm(&x, &self.norm2_weight, &self.norm2_bias)?;
let ffn_out = self.forward_ffn(&normed)?;
Ok((&x + &ffn_out)?)
}
fn forward_attn(&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)?.contiguous()?;
let k = k.reshape((n, h, hd))?.transpose(0, 1)?.contiguous()?;
let v = v.reshape((n, h, hd))?.transpose(0, 1)?.contiguous()?;
let scale = (hd as f64).sqrt();
let attn_weights = (q.matmul(&k.transpose(1, 2)?.contiguous()?)? / 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)?.contiguous()?.reshape((n, d))?;
Ok(attn_out
.matmul(&self.out_weight.t()?)?
.broadcast_add(&self.out_bias)?)
}
fn forward_ffn(&self, x: &Tensor) -> Result<Tensor> {
let h = x.matmul(&self.ffn_w1.t()?)?.broadcast_add(&self.ffn_b1)?;
let h = h.gelu_erf()?;
Ok(h.matmul(&self.ffn_w2.t()?)?.broadcast_add(&self.ffn_b2)?)
}
}
fn frozen_layer_norm(x: &Tensor, weight: &Tensor, bias: &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(weight)?.broadcast_add(bias)?)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MultiBranchTrainConfig {
pub output_dir: std::path::PathBuf,
pub epochs: usize,
pub batch_size: usize,
pub lr: f64,
pub weight_decay: f64,
pub patience: usize,
pub seed: u64,
pub min_lr: f64,
pub logit_adjust_tau: f64,
}
impl Default for MultiBranchTrainConfig {
fn default() -> Self {
Self {
output_dir: std::path::PathBuf::from("models/multi-branch-v1"),
epochs: 50,
batch_size: 64,
lr: 1e-3,
weight_decay: 0.01,
patience: 10,
seed: 42,
min_lr: 1e-6,
logit_adjust_tau: 0.0,
}
}
}
pub struct MultiBranchDataset {
pub char_feats: Vec<f32>,
pub embed_feats: Vec<f32>,
pub stats_feats: Vec<f32>,
pub header_feats: Vec<f32>,
pub valid_feats: Vec<f32>,
pub labels: Vec<u32>,
pub n_samples: usize,
pub char_dim: usize,
pub embed_dim: usize,
pub stats_dim: usize,
pub header_dim: usize,
pub valid_dim: usize,
pub table_groups: Vec<TableGroup>,
pub value_embeds: Vec<f32>,
pub value_mask: Vec<f32>,
pub n_values: usize,
pub value_embed_dim: usize,
}
impl MultiBranchDataset {
pub fn from_records(
records: &[TrainingRecord],
label_to_idx: &std::collections::HashMap<String, u32>,
char_dim: usize,
embed_dim: usize,
stats_dim: usize,
header_dim: usize,
) -> Result<Self> {
Self::from_records_with_groups(
records,
label_to_idx,
char_dim,
embed_dim,
stats_dim,
header_dim,
0,
None,
)
}
#[allow(clippy::too_many_arguments)]
pub fn from_records_with_groups(
records: &[TrainingRecord],
label_to_idx: &std::collections::HashMap<String, u32>,
char_dim: usize,
embed_dim: usize,
stats_dim: usize,
header_dim: usize,
valid_dim: usize,
table_groups: Option<Vec<TableGroup>>,
) -> Result<Self> {
let n = records.len();
let mut char_feats = Vec::with_capacity(n * char_dim);
let mut embed_feats_flat = Vec::with_capacity(n * embed_dim);
let mut stats_feats = Vec::with_capacity(n * stats_dim);
let mut header_feats = Vec::with_capacity(n * header_dim);
let mut valid_feats_flat = Vec::with_capacity(n * valid_dim);
let mut labels = Vec::with_capacity(n);
for record in records {
let idx = label_to_idx
.get(&record.label)
.copied()
.with_context(|| format!("Unknown label: {}", record.label))?;
labels.push(idx);
char_feats.extend_from_slice(&record.char_features);
embed_feats_flat.extend_from_slice(&record.embed_features);
stats_feats.extend_from_slice(&record.stats_features);
if record.header_features.len() >= header_dim {
header_feats.extend_from_slice(&record.header_features[..header_dim]);
} else {
header_feats.extend_from_slice(&record.header_features);
header_feats.extend(std::iter::repeat_n(
0.0f32,
header_dim - record.header_features.len(),
));
}
if valid_dim > 0 {
if record.validation_features.len() >= valid_dim {
valid_feats_flat.extend_from_slice(&record.validation_features[..valid_dim]);
} else {
valid_feats_flat.extend_from_slice(&record.validation_features);
valid_feats_flat.extend(std::iter::repeat_n(
0.0f32,
valid_dim - record.validation_features.len(),
));
}
}
}
let groups = table_groups.unwrap_or_else(|| {
if n > 0 {
vec![TableGroup {
record_indices: (0..n).collect(),
sibling_headers: Vec::new(),
}]
} else {
Vec::new()
}
});
Ok(Self {
char_feats,
embed_feats: embed_feats_flat,
stats_feats,
header_feats,
valid_feats: valid_feats_flat,
labels,
n_samples: n,
char_dim,
embed_dim,
stats_dim,
header_dim,
valid_dim,
table_groups: groups,
value_embeds: Vec::new(),
value_mask: Vec::new(),
n_values: 0,
value_embed_dim: 0,
})
}
pub fn with_value_attention(
mut self,
records: &[TrainingRecord],
cfg: &ValueAttentionConfig,
encoder: &Model2VecResources,
) -> Result<Self> {
let n_values = cfg.n_values;
let d = cfg.value_embed_dim;
let enc_dim = encoder.embed_dim()?;
if enc_dim != d {
bail!("value encoder dim {enc_dim} != config value_embed_dim {d}");
}
if records.len() != self.n_samples {
bail!(
"with_value_attention: {} records != dataset n_samples {}",
records.len(),
self.n_samples
);
}
let mut value_embeds = vec![0.0f32; self.n_samples * n_values * d];
let mut value_mask = vec![0.0f32; self.n_samples * n_values];
for (ri, record) in records.iter().enumerate() {
let take = record.values.len().min(n_values);
if take == 0 {
continue; }
let refs: Vec<&str> = record.values[..take].iter().map(|s| s.as_str()).collect();
let batch = encoder.encode_batch(&refs)?; for j in 0..take {
let row: Vec<f32> = batch.get(j)?.to_vec1()?;
let base = (ri * n_values + j) * d;
value_embeds[base..base + d].copy_from_slice(&row[..d]);
value_mask[ri * n_values + j] = 1.0;
}
}
self.value_embeds = value_embeds;
self.value_mask = value_mask;
self.n_values = n_values;
self.value_embed_dim = d;
Ok(self)
}
pub fn has_value_attention(&self) -> bool {
self.n_values > 0 && self.value_embed_dim > 0
}
pub fn expand_group_indices(&self, group_indices: &[usize]) -> Vec<usize> {
let mut out = Vec::new();
for &gi in group_indices {
out.extend_from_slice(&self.table_groups[gi].record_indices);
}
out
}
pub fn value_batch(
&self,
record_indices: &[usize],
device: &Device,
) -> candle_core::Result<Option<(Tensor, Tensor)>> {
if !self.has_value_attention() {
return Ok(None);
}
let bs = record_indices.len();
let n = self.n_values;
let d = self.value_embed_dim;
let mut embeds = Vec::with_capacity(bs * n * d);
let mut mask = Vec::with_capacity(bs * n);
for &i in record_indices {
let e_start = i * n * d;
embeds.extend_from_slice(&self.value_embeds[e_start..e_start + n * d]);
let m_start = i * n;
mask.extend_from_slice(&self.value_mask[m_start..m_start + n]);
}
let embeds_t = Tensor::new(embeds.as_slice(), device)?.reshape((bs, n, d))?;
let mask_t = Tensor::new(mask.as_slice(), device)?.reshape((bs, n))?;
Ok(Some((embeds_t, mask_t)))
}
pub fn len(&self) -> usize {
self.n_samples
}
pub fn is_empty(&self) -> bool {
self.n_samples == 0
}
#[allow(clippy::type_complexity)]
pub fn batch(
&self,
indices: &[usize],
device: &Device,
) -> candle_core::Result<(
Tensor,
Tensor,
Tensor,
Option<Tensor>,
Option<Tensor>,
Tensor,
)> {
let bs = indices.len();
let mut char_batch = Vec::with_capacity(bs * self.char_dim);
let mut embed_batch = Vec::with_capacity(bs * self.embed_dim);
let mut stats_batch = Vec::with_capacity(bs * self.stats_dim);
let mut header_batch = Vec::with_capacity(bs * self.header_dim);
let mut valid_batch = Vec::with_capacity(bs * self.valid_dim);
let mut label_batch = Vec::with_capacity(bs);
for &i in indices {
let char_start = i * self.char_dim;
char_batch.extend_from_slice(&self.char_feats[char_start..char_start + self.char_dim]);
let embed_start = i * self.embed_dim;
embed_batch
.extend_from_slice(&self.embed_feats[embed_start..embed_start + self.embed_dim]);
let stats_start = i * self.stats_dim;
stats_batch
.extend_from_slice(&self.stats_feats[stats_start..stats_start + self.stats_dim]);
if self.header_dim > 0 {
let header_start = i * self.header_dim;
header_batch.extend_from_slice(
&self.header_feats[header_start..header_start + self.header_dim],
);
}
if self.valid_dim > 0 {
let valid_start = i * self.valid_dim;
valid_batch.extend_from_slice(
&self.valid_feats[valid_start..valid_start + self.valid_dim],
);
}
label_batch.push(self.labels[i]);
}
let char_t = Tensor::new(char_batch.as_slice(), device)?.reshape((bs, self.char_dim))?;
let embed_t = Tensor::new(embed_batch.as_slice(), device)?.reshape((bs, self.embed_dim))?;
let stats_t = Tensor::new(stats_batch.as_slice(), device)?.reshape((bs, self.stats_dim))?;
let header_t = if self.header_dim > 0 {
Some(Tensor::new(header_batch.as_slice(), device)?.reshape((bs, self.header_dim))?)
} else {
None
};
let valid_t = if self.valid_dim > 0 {
Some(Tensor::new(valid_batch.as_slice(), device)?.reshape((bs, self.valid_dim))?)
} else {
None
};
let labels_t = Tensor::new(label_batch.as_slice(), device)?;
Ok((char_t, embed_t, stats_t, header_t, valid_t, labels_t))
}
#[allow(clippy::type_complexity)]
pub fn batch_groups(
&self,
group_indices: &[usize],
frozen_ctx: Option<&FrozenSiblingContext>,
device: &Device,
) -> Result<(
Tensor,
Tensor,
Tensor,
Option<Tensor>,
Option<Tensor>,
Tensor,
)> {
let mut all_indices = Vec::new();
let mut group_boundaries: Vec<(usize, usize, usize)> = Vec::new();
for &gi in group_indices {
let group = &self.table_groups[gi];
let start = all_indices.len();
all_indices.extend_from_slice(&group.record_indices);
let end = all_indices.len();
group_boundaries.push((start, end, gi));
}
let bs = all_indices.len();
if bs == 0 {
let char_t = Tensor::zeros((0, self.char_dim), DType::F32, device)?;
let embed_t = Tensor::zeros((0, self.embed_dim), DType::F32, device)?;
let stats_t = Tensor::zeros((0, self.stats_dim), DType::F32, device)?;
let header_t = if self.header_dim > 0 {
Some(Tensor::zeros((0, self.header_dim), DType::F32, device)?)
} else {
None
};
let valid_t = if self.valid_dim > 0 {
Some(Tensor::zeros((0, self.valid_dim), DType::F32, device)?)
} else {
None
};
let labels_t = Tensor::zeros(0, DType::U32, device)?;
return Ok((char_t, embed_t, stats_t, header_t, valid_t, labels_t));
}
let mut char_batch = Vec::with_capacity(bs * self.char_dim);
let mut embed_batch = Vec::with_capacity(bs * self.embed_dim);
let mut stats_batch = Vec::with_capacity(bs * self.stats_dim);
let mut header_batch = Vec::with_capacity(bs * self.header_dim);
let mut valid_batch = Vec::with_capacity(bs * self.valid_dim);
let mut label_batch = Vec::with_capacity(bs);
for &i in &all_indices {
let char_start = i * self.char_dim;
char_batch.extend_from_slice(&self.char_feats[char_start..char_start + self.char_dim]);
let embed_start = i * self.embed_dim;
embed_batch
.extend_from_slice(&self.embed_feats[embed_start..embed_start + self.embed_dim]);
let stats_start = i * self.stats_dim;
stats_batch
.extend_from_slice(&self.stats_feats[stats_start..stats_start + self.stats_dim]);
if self.header_dim > 0 {
let header_start = i * self.header_dim;
header_batch.extend_from_slice(
&self.header_feats[header_start..header_start + self.header_dim],
);
}
if self.valid_dim > 0 {
let valid_start = i * self.valid_dim;
valid_batch.extend_from_slice(
&self.valid_feats[valid_start..valid_start + self.valid_dim],
);
}
label_batch.push(self.labels[i]);
}
let char_t = Tensor::new(char_batch.as_slice(), device)?.reshape((bs, self.char_dim))?;
let embed_t = Tensor::new(embed_batch.as_slice(), device)?.reshape((bs, self.embed_dim))?;
let stats_t = Tensor::new(stats_batch.as_slice(), device)?.reshape((bs, self.stats_dim))?;
let labels_t = Tensor::new(label_batch.as_slice(), device)?;
let header_t = if self.header_dim > 0 {
let mut header_t =
Tensor::new(header_batch.as_slice(), device)?.reshape((bs, self.header_dim))?;
if let Some(ctx) = frozen_ctx {
for &(start, end, gi) in &group_boundaries {
let group = &self.table_groups[gi];
let n_cols = end - start;
if n_cols > 1 && !group.sibling_headers.is_empty() {
let group_headers = header_t.narrow(0, start, n_cols)?;
let enriched = ctx.forward(&group_headers)?;
let mut parts: Vec<Tensor> = Vec::new();
if start > 0 {
parts.push(header_t.narrow(0, 0, start)?);
}
parts.push(enriched);
if end < bs {
parts.push(header_t.narrow(0, end, bs - end)?);
}
header_t = Tensor::cat(&parts, 0)?;
}
}
}
Some(header_t)
} else {
None
};
let valid_t = if self.valid_dim > 0 {
Some(Tensor::new(valid_batch.as_slice(), device)?.reshape((bs, self.valid_dim))?)
} else {
None
};
Ok((char_t, embed_t, stats_t, header_t, valid_t, labels_t))
}
}
fn count_parameters(varmap: &VarMap) -> usize {
varmap
.all_vars()
.iter()
.map(|v| v.as_tensor().elem_count())
.sum()
}
fn compute_branch_gradient_norms(
varmap: &VarMap,
grads: &candle_core::backprop::GradStore,
) -> std::collections::HashMap<String, f32> {
let data = varmap.data().lock().unwrap();
let mut branch_sq_sums: std::collections::HashMap<String, f32> =
std::collections::HashMap::new();
for (name, var) in data.iter() {
let branch = if name.starts_with("char.") {
"char"
} else if name.starts_with("embed.") {
"embed"
} else if name.starts_with("stats.") {
"stats"
} else if name.starts_with("header.") {
"header"
} else if name.starts_with("valid.") {
"valid"
} else if name.starts_with("merge_") {
"merge"
} else if name.starts_with("head.") || name.starts_with("hier_") {
"head"
} else {
"other"
};
if let Some(grad) = grads.get(var.as_tensor()) {
if let Ok(sq_sum) = grad
.sqr()
.and_then(|sq| sq.sum_all())
.and_then(|s| s.to_dtype(candle_core::DType::F32))
.and_then(|s| s.to_scalar::<f32>())
{
*branch_sq_sums.entry(branch.to_string()).or_insert(0.0) += sq_sum;
}
}
}
branch_sq_sums
.into_iter()
.map(|(branch, sq_sum)| (branch, sq_sum.sqrt()))
.collect()
}
fn compute_hierarchical_loss(
domain_logits: &Tensor,
cat_logits_all: &[Tensor],
leaf_logits_all: &[Vec<Option<Tensor>>],
flat_labels: &Tensor,
hierarchy: &HierarchyMap,
device: &Device,
) -> candle_core::Result<Tensor> {
let flat_labels_vec: Vec<u32> = flat_labels.to_vec1()?;
let batch_len = flat_labels_vec.len();
let mut domain_targets = Vec::with_capacity(batch_len);
let mut cat_targets_by_domain: Vec<Vec<u32>> = vec![Vec::new(); hierarchy.num_domains()];
let mut cat_sample_indices_by_domain: Vec<Vec<usize>> =
vec![Vec::new(); hierarchy.num_domains()];
let mut leaf_targets_by_cat: Vec<Vec<Vec<u32>>> = Vec::new();
let mut leaf_sample_indices_by_cat: Vec<Vec<Vec<usize>>> = Vec::new();
for d in 0..hierarchy.num_domains() {
leaf_targets_by_cat.push(vec![Vec::new(); hierarchy.num_categories(d)]);
leaf_sample_indices_by_cat.push(vec![Vec::new(); hierarchy.num_categories(d)]);
}
for (i, &flat_idx) in flat_labels_vec.iter().enumerate() {
let (d, c, t) = hierarchy.flat_to_hier(flat_idx as usize);
domain_targets.push(d as u32);
cat_targets_by_domain[d].push(c as u32);
cat_sample_indices_by_domain[d].push(i);
leaf_targets_by_cat[d][c].push(t as u32);
leaf_sample_indices_by_cat[d][c].push(i);
}
let domain_target_tensor = Tensor::new(domain_targets, device)?;
let domain_loss = candle_nn::loss::cross_entropy(domain_logits, &domain_target_tensor)?;
let mut cat_loss_sum = Tensor::new(0.0f32, device)?;
let mut cat_count = 0usize;
for d in 0..hierarchy.num_domains() {
if cat_targets_by_domain[d].is_empty() {
continue;
}
let indices: Vec<u32> = cat_sample_indices_by_domain[d]
.iter()
.map(|&i| i as u32)
.collect();
let idx_tensor = Tensor::new(indices, device)?;
let cat_logits_subset = cat_logits_all[d].index_select(&idx_tensor, 0)?;
let cat_target_tensor = Tensor::new(cat_targets_by_domain[d].clone(), device)?;
let cl = candle_nn::loss::cross_entropy(&cat_logits_subset, &cat_target_tensor)?;
let n = cat_targets_by_domain[d].len() as f32;
cat_loss_sum = (cat_loss_sum + cl.broadcast_mul(&Tensor::new(n, device)?))?;
cat_count += cat_targets_by_domain[d].len();
}
let cat_loss = if cat_count > 0 {
cat_loss_sum.broadcast_div(&Tensor::new(cat_count as f32, device)?)?
} else {
Tensor::new(0.0f32, device)?
};
let mut leaf_loss_sum = Tensor::new(0.0f32, device)?;
let mut leaf_count = 0usize;
for d in 0..hierarchy.num_domains() {
for c in 0..hierarchy.num_categories(d) {
if hierarchy.is_degenerate(d, c) || leaf_targets_by_cat[d][c].is_empty() {
continue;
}
if let Some(ref leaf_logits) = leaf_logits_all[d][c] {
let indices: Vec<u32> = leaf_sample_indices_by_cat[d][c]
.iter()
.map(|&i| i as u32)
.collect();
let idx_tensor = Tensor::new(indices, device)?;
let leaf_logits_subset = leaf_logits.index_select(&idx_tensor, 0)?;
let leaf_target_tensor = Tensor::new(leaf_targets_by_cat[d][c].clone(), device)?;
let ll = candle_nn::loss::cross_entropy(&leaf_logits_subset, &leaf_target_tensor)?;
let n = leaf_targets_by_cat[d][c].len() as f32;
leaf_loss_sum = (leaf_loss_sum + ll.broadcast_mul(&Tensor::new(n, device)?))?;
leaf_count += leaf_targets_by_cat[d][c].len();
}
}
}
let leaf_loss = if leaf_count > 0 {
leaf_loss_sum.broadcast_div(&Tensor::new(leaf_count as f32, device)?)?
} else {
Tensor::new(0.0f32, device)?
};
let total = (domain_loss.broadcast_mul(&Tensor::new(0.2f32, device)?)?
+ cat_loss.broadcast_mul(&Tensor::new(0.3f32, device)?)?
+ leaf_loss.broadcast_mul(&Tensor::new(0.5f32, device)?)?)?;
Ok(total)
}
fn logit_adjust_prior(labels: &[u32], n_classes: usize, tau: f64) -> Vec<f32> {
let mut counts = vec![0u64; n_classes];
for &l in labels {
let idx = l as usize;
if idx < n_classes {
counts[idx] += 1;
}
}
let total: f64 = counts.iter().map(|&c| c as f64).sum::<f64>().max(1.0);
let eps = 1e-12_f64;
counts
.iter()
.map(|&c| (tau * ((c as f64 / total) + eps).ln()) as f32)
.collect()
}
pub fn train_multi_branch(
config: &MultiBranchTrainConfig,
model_config: &MultiBranchConfig,
train_data: &MultiBranchDataset,
val_data: &MultiBranchDataset,
labels: Option<&[String]>,
sibling_ctx_dir: Option<&Path>,
renderer: Option<Box<dyn crate::tui::TrainingRenderer>>,
) -> Result<crate::training::TrainingSummary> {
use crate::training::{
compute_accuracy, shuffled_batches, CosineScheduler, EarlyStopping, EpochMetrics,
};
use candle_nn::{AdamW, Optimizer, ParamsAdamW};
use rand::rngs::StdRng;
use rand::SeedableRng;
use std::collections::HashMap;
let (device, device_name) = crate::get_device();
eprintln!("Using {device_name} device");
let _seed_guard = crate::seeded_rng::seed_thread(config.seed);
let mut rng = StdRng::seed_from_u64(config.seed);
let frozen_ctx = match sibling_ctx_dir {
Some(dir) => match FrozenSiblingContext::load(dir, &device) {
Ok(ctx) => {
tracing::info!(
"Loaded frozen sibling-context attention ({}d) on training device",
ctx.embed_dim()
);
Some(ctx)
}
Err(e) => {
tracing::warn!("Failed to load sibling-context model: {e}");
None
}
},
None => None,
};
let frozen_ref = frozen_ctx.as_ref();
let is_hierarchical = model_config.head_type == HeadType::Hierarchical;
let use_group_batching = frozen_ref.is_some() && !train_data.table_groups.is_empty();
if frozen_ref.is_some() {
tracing::info!("Sibling-context enrichment enabled (frozen attention)");
}
tracing::info!(
"Starting multi-branch training: {} train, {} val, {} epochs, batch_size={}, lr={}, head={}",
train_data.len(),
val_data.len(),
config.epochs,
config.batch_size,
config.lr,
if is_hierarchical { "hierarchical" } else { "flat" },
);
let varmap = VarMap::new();
let vb = crate::seeded_rng::seeded_var_builder(&varmap, DType::F32, &device);
let model = if is_hierarchical {
let labels = labels
.ok_or_else(|| anyhow::anyhow!("Hierarchical head requires sorted labels list"))?;
MultiBranchModel::new_hierarchical(model_config, labels, vb)?
} else {
MultiBranchModel::new(model_config, vb)?
};
let n_params = count_parameters(&varmap);
tracing::info!("Model parameters: {}", n_params);
if model_config.has_header_branch() {
tracing::info!(
"Architecture: char [{} → {} → {}] | embed [{} → {} → {}] | stats [{} → {} → {}] | header [{} → {} → {}] | merge [{} → {} → {}] | head → {}",
model_config.char_dim, model_config.char_hidden[0], model_config.char_hidden[1],
model_config.embed_dim, model_config.embed_hidden[0], model_config.embed_hidden[1],
model_config.stats_dim, model_config.stats_hidden[0], model_config.stats_hidden[1],
model_config.header_dim, model_config.header_hidden[0], model_config.header_hidden[1],
model_config.merged_dim(),
model_config.merge_hidden[0], model_config.merge_hidden[1],
model_config.n_classes,
);
} else {
tracing::info!(
"Architecture: char [{} → {} → {}] | embed [{} → {} → {}] | stats [{} → {} → {}] | merge [{} → {} → {}] | head → {}",
model_config.char_dim, model_config.char_hidden[0], model_config.char_hidden[1],
model_config.embed_dim, model_config.embed_hidden[0], model_config.embed_hidden[1],
model_config.stats_dim, model_config.stats_hidden[0], model_config.stats_hidden[1],
model_config.merged_dim(),
model_config.merge_hidden[0], model_config.merge_hidden[1],
model_config.n_classes,
);
}
if model_config.has_validation_branch() {
tracing::info!(
"Validation branch: [{} → {} → {}]",
model_config.valid_dim,
model_config.valid_hidden[0],
model_config.valid_hidden[1],
);
}
if is_hierarchical {
let hier = model.hierarchical_head().unwrap();
let h = hier.hierarchy();
tracing::info!(
"Hierarchical: {} domains, {} categories, {} degenerate",
h.num_domains(),
h.total_categories(),
(0..h.num_domains())
.flat_map(|d| (0..h.num_categories(d)).map(move |c| (d, c)))
.filter(|&(d, c)| h.is_degenerate(d, c))
.count()
);
}
let adamw_params = ParamsAdamW {
lr: config.lr,
weight_decay: config.weight_decay,
..Default::default()
};
let mut optimizer = AdamW::new(varmap.all_vars(), adamw_params)?;
let logit_adjust: Option<Tensor> = if config.logit_adjust_tau > 0.0 && !is_hierarchical {
let n = model_config.n_classes;
let adj = logit_adjust_prior(&train_data.labels, n, config.logit_adjust_tau);
tracing::info!(
"Logit adjustment ENABLED (τ={:.3}): rare classes up-weighted via train-time logit prior (choice 0097)",
config.logit_adjust_tau
);
Some(Tensor::from_vec(adj, (1, n), &device)?)
} else {
None
};
let scheduler = CosineScheduler::new(config.lr, config.min_lr, config.epochs);
let mut early_stopping = EarlyStopping::new(config.patience, true);
std::fs::create_dir_all(&config.output_dir).with_context(|| {
format!(
"Failed to create output dir: {}",
config.output_dir.display()
)
})?;
model_config.save(&config.output_dir.join("config.json"))?;
let mut epoch_metrics = Vec::new();
let total_start = std::time::Instant::now();
let mut renderer: Box<dyn crate::tui::TrainingRenderer> =
renderer.unwrap_or_else(|| Box::new(crate::tui::LogRenderer::new()));
let initial_batch_count = train_data.len().div_ceil(config.batch_size);
renderer.on_train_start(config.epochs, initial_batch_count);
for epoch in 0..config.epochs {
let epoch_start = std::time::Instant::now();
let lr = scheduler.lr(epoch);
optimizer.set_learning_rate(lr);
let mut train_loss_sum = 0.0f64;
let mut train_correct_sum = 0.0f64;
let mut train_samples = 0usize;
let mut epoch_grad_norms: HashMap<String, f64> = HashMap::new();
let mut grad_norm_batch_count = 0usize;
let group_batches: Vec<Vec<usize>>;
let flat_batches: Vec<Vec<usize>>;
let total_batches;
if use_group_batching {
use rand::seq::SliceRandom;
let n_groups = train_data.table_groups.len();
let mut group_order: Vec<usize> = (0..n_groups).collect();
group_order.shuffle(&mut rng);
let mut batches = Vec::new();
let mut current_batch = Vec::new();
let mut current_size = 0usize;
for &gi in &group_order {
let group_size = train_data.table_groups[gi].record_indices.len();
current_batch.push(gi);
current_size += group_size;
if current_size >= config.batch_size {
batches.push(current_batch);
current_batch = Vec::new();
current_size = 0;
}
}
if !current_batch.is_empty() {
batches.push(current_batch);
}
group_batches = batches;
flat_batches = Vec::new();
total_batches = group_batches.len();
} else {
group_batches = Vec::new();
flat_batches = shuffled_batches(train_data.len(), config.batch_size, &mut rng);
total_batches = flat_batches.len();
}
for batch_num in 0..total_batches {
let record_idx_list: Vec<usize> = if use_group_batching {
train_data.expand_group_indices(&group_batches[batch_num])
} else {
flat_batches[batch_num].clone()
};
let (char_t, embed_t, stats_t, header_t, valid_t, labels_t) = if use_group_batching {
train_data.batch_groups(&group_batches[batch_num], frozen_ref, &device)?
} else {
train_data.batch(&flat_batches[batch_num], &device)?
};
let bs = labels_t.dim(0)?;
let value_bt = train_data.value_batch(&record_idx_list, &device)?;
let (ve, vm) = match &value_bt {
Some((e, m)) => (Some(e), Some(m)),
None => (None, None),
};
let embed_train = model.embed_input(&embed_t, ve, vm, true)?;
let loss = if is_hierarchical {
let hier = model.hierarchical_head().unwrap();
let hierarchy = hier.hierarchy();
let (domain_logits, cat_logits, leaf_logits) = model.forward_levels(
&char_t,
&embed_train,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
true,
)?;
compute_hierarchical_loss(
&domain_logits,
&cat_logits,
&leaf_logits,
&labels_t,
hierarchy,
&device,
)?
} else {
let logits = model.forward(
&char_t,
&embed_train,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
true,
)?;
match logit_adjust {
Some(ref adj) => {
candle_nn::loss::cross_entropy(&logits.broadcast_add(adj)?, &labels_t)?
}
None => candle_nn::loss::cross_entropy(&logits, &labels_t)?,
}
};
let grads = loss.backward()?;
let batch_norms = compute_branch_gradient_norms(&varmap, &grads);
for (branch, norm) in &batch_norms {
*epoch_grad_norms.entry(branch.clone()).or_insert(0.0) += *norm as f64;
}
grad_norm_batch_count += 1;
optimizer.step(&grads)?;
let loss_val: f32 = loss.to_scalar()?;
train_loss_sum += loss_val as f64 * bs as f64;
let embed_eval = model.embed_input(&embed_t, ve, vm, false)?;
let output = model.forward(
&char_t,
&embed_eval,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
false,
)?;
let acc = compute_accuracy(&output, &labels_t)?;
train_correct_sum += acc as f64 * bs as f64;
train_samples += bs;
renderer.on_batch_end(epoch, batch_num + 1, total_batches, loss_val);
}
let train_loss = (train_loss_sum / train_samples as f64) as f32;
let train_accuracy = (train_correct_sum / train_samples as f64) as f32;
let (val_accuracy, val_loss) = {
let val_indices: Vec<usize> = (0..val_data.len()).collect();
let val_batches: Vec<Vec<usize>> = val_indices
.chunks(config.batch_size)
.map(|c| c.to_vec())
.collect();
let mut val_loss_sum = 0.0f64;
let mut val_correct_sum = 0.0f64;
let mut val_samples = 0usize;
for batch_idx in &val_batches {
let (char_t, embed_t, stats_t, header_t, valid_t, labels_t) =
val_data.batch(batch_idx, &device)?;
let bs = batch_idx.len();
let value_bt = val_data.value_batch(batch_idx, &device)?;
let (ve, vm) = match &value_bt {
Some((e, m)) => (Some(e), Some(m)),
None => (None, None),
};
let embed_t = model.embed_input(&embed_t, ve, vm, false)?;
let loss_val = if is_hierarchical {
let hier = model.hierarchical_head().unwrap();
let hierarchy = hier.hierarchy();
let (domain_logits, cat_logits, leaf_logits) = model.forward_levels(
&char_t,
&embed_t,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
false,
)?;
let loss = compute_hierarchical_loss(
&domain_logits,
&cat_logits,
&leaf_logits,
&labels_t,
hierarchy,
&device,
)?;
loss.to_scalar::<f32>()?
} else {
let logits = model.forward(
&char_t,
&embed_t,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
false,
)?;
let loss = candle_nn::loss::cross_entropy(&logits, &labels_t)?;
loss.to_scalar::<f32>()?
};
val_loss_sum += loss_val as f64 * bs as f64;
let output = model.forward(
&char_t,
&embed_t,
&stats_t,
header_t.as_ref(),
valid_t.as_ref(),
false,
)?;
let acc = compute_accuracy(&output, &labels_t)?;
val_correct_sum += acc as f64 * bs as f64;
val_samples += bs;
}
let val_acc = (val_correct_sum / val_samples as f64) as f32;
let val_loss = (val_loss_sum / val_samples as f64) as f32;
(val_acc, val_loss)
};
let epoch_time = epoch_start.elapsed().as_secs_f32();
let branch_gradient_norms = if grad_norm_batch_count > 0 {
Some(
epoch_grad_norms
.iter()
.map(|(branch, sum)| {
(branch.clone(), (*sum / grad_norm_batch_count as f64) as f32)
})
.collect::<HashMap<String, f32>>(),
)
} else {
None
};
let metrics = EpochMetrics {
epoch,
train_loss,
val_loss,
train_accuracy,
val_accuracy,
learning_rate: lr,
epoch_time_secs: epoch_time,
branch_gradient_norms,
};
renderer.on_epoch_end(&metrics);
epoch_metrics.push(metrics);
tracing::info!(
"Epoch {:>3}/{}: train_loss={:.4} val_loss={:.4} train_acc={:.3} val_acc={:.3} lr={:.2e} ({:.1}s)",
epoch + 1,
config.epochs,
train_loss,
val_loss,
train_accuracy,
val_accuracy,
lr,
epoch_time,
);
let results_path = config.output_dir.join("results.json");
let results_tmp = config.output_dir.join("results.json.tmp");
if let Ok(json) = serde_json::to_string_pretty(&epoch_metrics) {
let _ = std::fs::write(&results_tmp, &json)
.and_then(|_| std::fs::rename(&results_tmp, &results_path));
}
let epochs_path = config.output_dir.join("epochs.jsonl");
if let Ok(mut f) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(&epochs_path)
{
let _ = writeln!(
f,
r#"{{"epoch":{},"train_loss":{:.4},"val_loss":{:.4},"train_acc":{:.4},"val_acc":{:.4},"lr":{:.2e},"time":{:.1}}}"#,
epoch + 1,
train_loss,
val_loss,
train_accuracy,
val_accuracy,
lr,
epoch_time,
);
}
let should_stop = early_stopping.step(epoch, val_accuracy);
if early_stopping.best_epoch() == epoch {
let checkpoint_path = config.output_dir.join("model_best.safetensors");
varmap.save(&checkpoint_path)?;
tracing::info!(" -> New best model saved (val_acc={:.3})", val_accuracy);
}
if should_stop {
tracing::info!(
"Early stopping at epoch {} (best epoch {})",
epoch + 1,
early_stopping.best_epoch() + 1,
);
break;
}
}
renderer.on_train_end();
let total_time = total_start.elapsed().as_secs_f32();
let best_path = config.output_dir.join("model_best.safetensors");
let final_path = config.output_dir.join("model.safetensors");
if best_path.exists() && best_path != final_path {
std::fs::rename(&best_path, &final_path)?;
}
model_config.save(&config.output_dir.join("config.json"))?;
let results_json = serde_json::to_string_pretty(&epoch_metrics)?;
std::fs::write(config.output_dir.join("results.json"), &results_json)?;
let total_epochs = epoch_metrics.len();
tracing::info!(
"Training complete: best_epoch={}, val_acc={:.3}, {:.1}s total",
early_stopping.best_epoch() + 1,
early_stopping.best_metric(),
total_time,
);
Ok(crate::training::TrainingSummary {
best_epoch: early_stopping.best_epoch(),
best_val_accuracy: early_stopping.best_metric(),
total_epochs,
total_time_secs: total_time,
epoch_metrics,
})
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::{DType, Device};
use candle_nn::{Optimizer, VarMap};
fn make_config() -> MultiBranchConfig {
MultiBranchConfig::default()
}
#[test]
fn test_value_attention_end_to_end() {
let enc_dir = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.parent()
.unwrap()
.parent()
.unwrap()
.join("models")
.join("model2vec");
if !enc_dir.join("model.safetensors").exists() {
eprintln!("skip test_value_attention_end_to_end: models/model2vec absent");
return;
}
let enc = Model2VecResources::load(&enc_dir).unwrap();
let d = enc.embed_dim().unwrap();
let va = ValueAttentionConfig {
value_embed_dim: d,
n_values: 8,
n_heads: 4,
self_attn_layers: 1,
ffn_hidden: 64,
pool_slots: 2,
dropout: 0.0,
keep_blender_concat: true,
};
let blender = 16usize;
let config = MultiBranchConfig {
char_dim: 12,
embed_dim: blender,
stats_dim: 6,
header_dim: 0,
char_hidden: [16, 8],
embed_hidden: [16, 8],
stats_hidden: [8, 4],
header_hidden: [0, 0],
valid_dim: 0,
valid_hidden: [0, 0],
merge_hidden: [16, 16],
n_classes: 3,
dropout: 0.0,
head_type: HeadType::Flat,
activation: Activation::ReLU,
use_layer_norm: false,
value_attention: Some(va.clone()),
};
let label_names = ["a.b.c0", "a.b.c1", "a.b.c2"];
let value_sets = [
vec!["alice@example.com", "bob@example.com"],
vec!["2024-01-15", "2023-06-30", "2022-12-01"],
vec!["London", "Paris"],
];
let records: Vec<TrainingRecord> = (0..6)
.map(|i| {
let c = i % 3;
TrainingRecord {
label: label_names[c].to_string(),
char_features: vec![c as f32 * 0.1; 12],
embed_features: vec![c as f32 * 0.2; blender],
stats_features: vec![c as f32 * 0.3; 6],
header_features: vec![],
validation_features: vec![],
values: value_sets[c].iter().map(|s| s.to_string()).collect(),
}
})
.collect();
let mut label_to_idx = std::collections::HashMap::new();
for (i, n) in label_names.iter().enumerate() {
label_to_idx.insert(n.to_string(), i as u32);
}
let ds = MultiBranchDataset::from_records(&records, &label_to_idx, 12, blender, 6, 0)
.unwrap()
.with_value_attention(&records, &va, &enc)
.unwrap();
assert!(ds.has_value_attention());
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
assert!(model.has_value_attention());
let indices: Vec<usize> = (0..records.len()).collect();
let (char_t, embed_t, stats_t, _h, _v, labels_t) = ds.batch(&indices, &device).unwrap();
let (ve, vm) = ds.value_batch(&indices, &device).unwrap().unwrap();
let embed_in = model
.embed_input(&embed_t, Some(&ve), Some(&vm), true)
.unwrap();
assert_eq!(embed_in.dim(1).unwrap(), blender + va.output_dim());
let vars = varmap.all_vars();
let before: Vec<Vec<f32>> = vars
.iter()
.map(|v| v.flatten_all().unwrap().to_vec1::<f32>().unwrap())
.collect();
let adamw_params = candle_nn::ParamsAdamW {
lr: 0.01,
..Default::default()
};
let mut opt = candle_nn::AdamW::new(varmap.all_vars(), adamw_params).unwrap();
let logits = model
.forward(&char_t, &embed_in, &stats_t, None, None, true)
.unwrap();
assert_eq!(logits.dims(), &[6, 3]);
let loss = candle_nn::loss::cross_entropy(&logits, &labels_t).unwrap();
assert!(loss.to_scalar::<f32>().unwrap().is_finite());
let grads = loss.backward().unwrap();
opt.step(&grads).unwrap();
let changed = vars
.iter()
.zip(before.iter())
.filter(|(v, b)| &v.flatten_all().unwrap().to_vec1::<f32>().unwrap() != *b)
.count();
assert!(
changed > 0,
"no parameters updated by the value-attention step"
);
}
#[test]
fn test_logit_adjust_prior_upweights_rare_classes() {
let labels = vec![0u32, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2];
let tau = 1.0;
let adj = logit_adjust_prior(&labels, 3, tau);
assert_eq!(adj.len(), 3);
assert!(adj[2] < adj[1], "rarer class must have more-negative prior");
assert!(adj[1] < adj[0], "rarer class must have more-negative prior");
let expected0 = (8.0_f64 / 13.0).ln() as f32;
assert!((adj[0] - expected0).abs() < 1e-5);
let adj_half = logit_adjust_prior(&labels, 3, 0.5);
assert!((adj_half[2] - adj[2] * 0.5).abs() < 1e-5);
}
#[test]
fn test_logit_adjust_prior_handles_empty_class() {
let labels = vec![0u32, 0, 1];
let adj = logit_adjust_prior(&labels, 3, 1.0);
assert!(adj[2].is_finite(), "empty class prior must be finite");
assert!(adj[2] < adj[1] && adj[2] < adj[0]);
}
#[test]
fn test_forward_pass_shape() {
let config = make_config();
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 10;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let logits = model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, true)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
let logits = model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, false)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
}
#[test]
fn test_gradient_flow() {
let config = make_config();
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 4;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let targets = Tensor::new(&[0u32, 1, 2, 3], &device).unwrap();
let vars = varmap.all_vars();
let initial_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let logits = model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, true)
.unwrap();
let loss = candle_nn::loss::cross_entropy(&logits, &targets).unwrap();
let adamw_params = candle_nn::ParamsAdamW {
lr: 0.01,
..Default::default()
};
let mut optimizer = candle_nn::AdamW::new(varmap.all_vars(), adamw_params).unwrap();
optimizer.backward_step(&loss).unwrap();
let updated_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let mut any_changed = false;
for (initial, updated) in initial_values.iter().zip(updated_values.iter()) {
if initial != updated {
any_changed = true;
break;
}
}
assert!(
any_changed,
"At least some parameters should have changed after backward pass"
);
let n_changed: usize = initial_values
.iter()
.zip(updated_values.iter())
.filter(|(a, b)| a != b)
.count();
assert!(
n_changed > 5,
"Expected many parameters to change, only {} changed out of {}",
n_changed,
vars.len()
);
}
#[test]
fn test_config_serialization() {
let config = MultiBranchConfig {
char_dim: 960,
embed_dim: 512,
stats_dim: 27,
header_dim: 128,
valid_dim: 0,
char_hidden: [300, 300],
embed_hidden: [200, 200],
stats_hidden: [128, 64],
header_hidden: [128, 64],
valid_hidden: [0, 0],
merge_hidden: [500, 500],
n_classes: 250,
dropout: 0.35,
head_type: HeadType::Flat,
activation: Activation::ReLU,
use_layer_norm: false,
value_attention: None,
};
let tmp = tempfile::NamedTempFile::new().unwrap();
config.save(tmp.path()).unwrap();
let loaded = MultiBranchConfig::load(tmp.path()).unwrap();
assert_eq!(config.char_dim, loaded.char_dim);
assert_eq!(config.embed_dim, loaded.embed_dim);
assert_eq!(config.stats_dim, loaded.stats_dim);
assert_eq!(config.header_dim, loaded.header_dim);
assert_eq!(config.char_hidden, loaded.char_hidden);
assert_eq!(config.embed_hidden, loaded.embed_hidden);
assert_eq!(config.stats_hidden, loaded.stats_hidden);
assert_eq!(config.merge_hidden, loaded.merge_hidden);
assert_eq!(config.n_classes, loaded.n_classes);
assert!((config.dropout - loaded.dropout).abs() < 1e-6);
assert_eq!(config.head_type, loaded.head_type);
assert_eq!(config.activation, loaded.activation);
assert_eq!(config.use_layer_norm, loaded.use_layer_norm);
}
#[test]
fn test_config_serialization_gelu_layer_norm() {
let config = MultiBranchConfig {
activation: Activation::GELU,
use_layer_norm: true,
..MultiBranchConfig::default()
};
let tmp = tempfile::NamedTempFile::new().unwrap();
config.save(tmp.path()).unwrap();
let loaded = MultiBranchConfig::load(tmp.path()).unwrap();
assert_eq!(loaded.activation, Activation::GELU);
assert!(loaded.use_layer_norm);
}
#[test]
fn test_config_backward_compat_deserializes_without_new_fields() {
let json = r#"{
"char_dim": 960, "embed_dim": 512, "stats_dim": 27, "header_dim": 128,
"char_hidden": [300, 300], "embed_hidden": [200, 200],
"stats_hidden": [128, 64], "header_hidden": [128, 64],
"merge_hidden": [500, 500], "n_classes": 250, "dropout": 0.35,
"head_type": "Flat"
}"#;
let config: MultiBranchConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.activation, Activation::ReLU);
assert!(!config.use_layer_norm);
}
#[test]
fn test_forward_pass_shape_gelu_layer_norm() {
let config = MultiBranchConfig {
activation: Activation::GELU,
use_layer_norm: true,
..MultiBranchConfig::default()
};
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 10;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
None,
true,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
None,
false,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
}
#[test]
fn test_gradient_flow_gelu_layer_norm() {
let config = MultiBranchConfig {
activation: Activation::GELU,
use_layer_norm: true,
..MultiBranchConfig::default()
};
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 4;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let targets = Tensor::new(&[0u32, 1, 2, 3], &device).unwrap();
let vars = varmap.all_vars();
let initial_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
None,
true,
)
.unwrap();
let loss = candle_nn::loss::cross_entropy(&logits, &targets).unwrap();
let adamw_params = candle_nn::ParamsAdamW {
lr: 0.01,
..Default::default()
};
let mut optimizer = candle_nn::AdamW::new(varmap.all_vars(), adamw_params).unwrap();
optimizer.backward_step(&loss).unwrap();
let updated_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let n_changed: usize = initial_values
.iter()
.zip(updated_values.iter())
.filter(|(a, b)| a != b)
.count();
assert!(
n_changed > 5,
"Expected many parameters to change with GELU+LN, only {} changed out of {}",
n_changed,
vars.len()
);
}
#[test]
fn test_gelu_vs_relu_outputs_differ() {
let device = Device::Cpu;
let relu_config = MultiBranchConfig::default(); let gelu_config = MultiBranchConfig {
activation: Activation::GELU,
use_layer_norm: true,
..MultiBranchConfig::default()
};
let varmap_relu = VarMap::new();
let vb_relu = VarBuilder::from_varmap(&varmap_relu, DType::F32, &device);
let relu_model = MultiBranchModel::new(&relu_config, vb_relu).unwrap();
let varmap_gelu = VarMap::new();
let vb_gelu = VarBuilder::from_varmap(&varmap_gelu, DType::F32, &device);
let gelu_model = MultiBranchModel::new(&gelu_config, vb_gelu).unwrap();
let char_feats = Tensor::randn(0.0f32, 1.0, (1, 960), &device).unwrap();
let embed_feats = Tensor::randn(0.0f32, 1.0, (1, 512), &device).unwrap();
let stats_feats = Tensor::randn(0.0f32, 1.0, (1, 27), &device).unwrap();
let relu_out = relu_model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, false)
.unwrap();
let gelu_out = gelu_model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, false)
.unwrap();
assert_eq!(relu_out.dims(), gelu_out.dims());
let relu_vec = relu_out.flatten_all().unwrap().to_vec1::<f32>().unwrap();
let gelu_vec = gelu_out.flatten_all().unwrap().to_vec1::<f32>().unwrap();
let any_different = relu_vec
.iter()
.zip(gelu_vec.iter())
.any(|(a, b)| (a - b).abs() > 1e-6);
assert!(
any_different,
"ReLU and GELU models should produce different outputs"
);
}
#[test]
fn test_training_data_roundtrip() {
let records: Vec<TrainingRecord> = (0..10)
.map(|i| TrainingRecord {
label: format!("identity.person.type_{}", i),
char_features: (0..960).map(|j| (i * 960 + j) as f32 * 0.001).collect(),
embed_features: (0..512).map(|j| (i * 512 + j) as f32 * 0.002).collect(),
stats_features: (0..27).map(|j| (i * 27 + j) as f32 * 0.1).collect(),
header_features: (0..128).map(|j| (i * 128 + j) as f32 * 0.003).collect(),
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data(tmp.path(), &records, 960, 512, 27, 128).unwrap();
let (header, loaded, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.n_records, 10);
assert_eq!(header.version, 2);
assert_eq!(header.char_dim, 960);
assert_eq!(header.embed_dim, 512);
assert_eq!(header.stats_dim, 27);
assert_eq!(header.header_dim, 128);
assert_eq!(loaded.len(), 10);
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features.len(), read.char_features.len());
assert_eq!(orig.embed_features.len(), read.embed_features.len());
assert_eq!(orig.stats_features.len(), read.stats_features.len());
assert_eq!(orig.header_features.len(), read.header_features.len());
for (a, b) in orig.char_features.iter().zip(read.char_features.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "char feature mismatch");
}
for (a, b) in orig.embed_features.iter().zip(read.embed_features.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "embed feature mismatch");
}
for (a, b) in orig.stats_features.iter().zip(read.stats_features.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "stats feature mismatch");
}
for (a, b) in orig.header_features.iter().zip(read.header_features.iter()) {
assert_eq!(a.to_bits(), b.to_bits(), "header feature mismatch");
}
}
}
#[test]
fn test_ftmb_v1_read_compat() {
use std::io::Write;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mut file = std::fs::File::create(tmp.path()).unwrap();
let char_dim: u16 = 4;
let embed_dim: u16 = 3;
let stats_dim: u16 = 2;
let n_records: u64 = 2;
file.write_all(b"FTMB").unwrap();
file.write_all(&1u32.to_le_bytes()).unwrap(); file.write_all(&n_records.to_le_bytes()).unwrap();
file.write_all(&char_dim.to_le_bytes()).unwrap();
file.write_all(&embed_dim.to_le_bytes()).unwrap();
file.write_all(&stats_dim.to_le_bytes()).unwrap();
file.write_all(&[0u8; 2]).unwrap();
for i in 0..2 {
let label = format!("type_{i}");
let label_bytes = label.as_bytes();
file.write_all(&(label_bytes.len() as u16).to_le_bytes())
.unwrap();
file.write_all(label_bytes).unwrap();
for j in 0..char_dim {
file.write_all(&((i * 10 + j) as f32).to_le_bytes())
.unwrap();
}
for j in 0..embed_dim {
file.write_all(&((i * 10 + j) as f32 * 0.5).to_le_bytes())
.unwrap();
}
for j in 0..stats_dim {
file.write_all(&((i * 10 + j) as f32 * 0.1).to_le_bytes())
.unwrap();
}
}
drop(file);
let (header, records, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 1);
assert_eq!(header.n_records, 2);
assert_eq!(header.char_dim, 4);
assert_eq!(header.embed_dim, 3);
assert_eq!(header.stats_dim, 2);
assert_eq!(header.header_dim, 0);
assert_eq!(records.len(), 2);
for record in &records {
assert!(
record.header_features.is_empty(),
"v1 should have empty header_features"
);
assert_eq!(record.char_features.len(), 4);
assert_eq!(record.embed_features.len(), 3);
assert_eq!(record.stats_features.len(), 2);
}
assert_eq!(records[0].label, "type_0");
assert_eq!(records[1].label, "type_1");
}
#[test]
fn test_ftmb_v2_roundtrip_4_vectors() {
let records: Vec<TrainingRecord> = (0..3)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32 * 1.0; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data(tmp.path(), &records, 8, 6, 4, 5).unwrap();
let (header, loaded, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 2);
assert_eq!(header.n_records, 3);
assert_eq!(header.char_dim, 8);
assert_eq!(header.embed_dim, 6);
assert_eq!(header.stats_dim, 4);
assert_eq!(header.header_dim, 5);
assert_eq!(loaded.len(), 3);
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.embed_features, read.embed_features);
assert_eq!(orig.stats_features, read.stats_features);
assert_eq!(orig.header_features, read.header_features);
}
}
#[test]
fn test_dataset_batch() {
let records: Vec<TrainingRecord> = (0..5)
.map(|i| TrainingRecord {
label: format!("identity.person.type_{}", i % 3),
char_features: vec![i as f32; 960],
embed_features: vec![i as f32 * 2.0; 512],
stats_features: vec![i as f32 * 3.0; 27],
header_features: vec![i as f32 * 4.0; 128],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let mut label_to_idx = std::collections::HashMap::new();
label_to_idx.insert("identity.person.type_0".to_string(), 0u32);
label_to_idx.insert("identity.person.type_1".to_string(), 1u32);
label_to_idx.insert("identity.person.type_2".to_string(), 2u32);
let dataset =
MultiBranchDataset::from_records(&records, &label_to_idx, 960, 512, 27, 128).unwrap();
assert_eq!(dataset.len(), 5);
let device = Device::Cpu;
let (char_t, embed_t, stats_t, header_t, _valid_t, labels_t) =
dataset.batch(&[0, 2, 4], &device).unwrap();
assert_eq!(char_t.dims(), &[3, 960]);
assert_eq!(embed_t.dims(), &[3, 512]);
assert_eq!(stats_t.dims(), &[3, 27]);
let header_t = header_t.expect("header tensor should be Some for header_dim=128");
assert_eq!(header_t.dims(), &[3, 128]);
assert_eq!(labels_t.dims(), &[3]);
let labels: Vec<u32> = labels_t.to_vec1().unwrap();
assert_eq!(labels, vec![0, 2, 1]); }
#[test]
fn test_training_loop_small() {
let config = MultiBranchConfig {
n_classes: 3,
..Default::default()
};
let mut records = Vec::new();
let labels = ["type_a", "type_b", "type_c"];
for i in 0..30 {
let class_idx = i % 3;
let bias = class_idx as f32;
records.push(TrainingRecord {
label: labels[class_idx].to_string(),
char_features: (0..960)
.map(|j| if j % 3 == class_idx { 1.0 + bias } else { 0.1 })
.collect(),
embed_features: (0..512)
.map(|j| if j % 3 == class_idx { 1.0 + bias } else { 0.1 })
.collect(),
stats_features: (0..27)
.map(|j| if j % 3 == class_idx { 1.0 + bias } else { 0.1 })
.collect(),
header_features: (0..128)
.map(|j| if j % 3 == class_idx { 0.5 + bias } else { 0.05 })
.collect(),
validation_features: Vec::new(),
values: Vec::new(),
});
}
let mut label_to_idx = std::collections::HashMap::new();
label_to_idx.insert("type_a".to_string(), 0u32);
label_to_idx.insert("type_b".to_string(), 1u32);
label_to_idx.insert("type_c".to_string(), 2u32);
let train_data =
MultiBranchDataset::from_records(&records[..20], &label_to_idx, 960, 512, 27, 128)
.unwrap();
let val_data =
MultiBranchDataset::from_records(&records[20..], &label_to_idx, 960, 512, 27, 128)
.unwrap();
let tmp_dir = tempfile::tempdir().unwrap();
let train_config = MultiBranchTrainConfig {
output_dir: tmp_dir.path().to_path_buf(),
epochs: 5,
batch_size: 10,
lr: 1e-3,
weight_decay: 1e-4,
patience: 10,
seed: 42,
min_lr: 1e-6,
logit_adjust_tau: 0.0,
};
let summary = train_multi_branch(
&train_config,
&config,
&train_data,
&val_data,
None,
None,
None,
)
.unwrap();
assert_eq!(summary.total_epochs, 5);
assert_eq!(summary.epoch_metrics.len(), 5);
let loss_0 = summary.epoch_metrics[0].train_loss;
let loss_4 = summary.epoch_metrics[4].train_loss;
assert!(
loss_4 < loss_0,
"Training loss should decrease: epoch 0 = {}, epoch 4 = {}",
loss_0,
loss_4,
);
assert!(tmp_dir.path().join("model.safetensors").exists());
assert!(tmp_dir.path().join("config.json").exists());
assert!(tmp_dir.path().join("results.json").exists());
}
#[test]
fn test_merged_dim() {
let config = MultiBranchConfig::default();
assert_eq!(config.merged_dim(), 300 + 200 + 64 + 64);
let old_config = MultiBranchConfig {
header_dim: 0,
header_hidden: [0, 0],
..Default::default()
};
assert_eq!(old_config.merged_dim(), 300 + 200 + 64); }
fn make_hier_labels() -> Vec<String> {
vec![
"container.array.comma_separated".to_string(),
"container.array.pipe_separated".to_string(),
"container.object.json".to_string(),
"datetime.date.iso".to_string(),
"datetime.date.ymd_slash".to_string(),
"datetime.time.hms_24h".to_string(),
"geography.location.city".to_string(),
"geography.location.country".to_string(),
"identity.person.email".to_string(),
"identity.person.full_name".to_string(),
]
}
#[test]
fn test_hierarchical_forward_pass_shape() {
let labels = make_hier_labels();
let config = MultiBranchConfig {
n_classes: labels.len(),
..Default::default()
};
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new_hierarchical(&config, &labels, vb).unwrap();
assert!(model.is_hierarchical());
let batch_size = 5;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let output = model
.forward(&char_feats, &embed_feats, &stats_feats, None, None, false)
.unwrap();
assert_eq!(output.dims(), &[batch_size, labels.len()]);
let (domain_logits, cat_logits, leaf_logits) = model
.forward_levels(&char_feats, &embed_feats, &stats_feats, None, None, false)
.unwrap();
assert_eq!(domain_logits.dims()[0], batch_size);
assert_eq!(domain_logits.dims()[1], 4);
assert_eq!(cat_logits.len(), 4);
assert_eq!(leaf_logits.len(), 4);
}
#[test]
fn test_hierarchical_gradient_flow() {
let labels = make_hier_labels();
let config = MultiBranchConfig {
n_classes: labels.len(),
..Default::default()
};
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new_hierarchical(&config, &labels, vb).unwrap();
let batch_size = 4;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let targets = Tensor::new(&[0u32, 3, 6, 8], &device).unwrap();
let vars = varmap.all_vars();
let initial_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let hier = model.hierarchical_head().unwrap();
let hierarchy = hier.hierarchy();
let (domain_logits, cat_logits, leaf_logits) = model
.forward_levels(&char_feats, &embed_feats, &stats_feats, None, None, true)
.unwrap();
let loss = compute_hierarchical_loss(
&domain_logits,
&cat_logits,
&leaf_logits,
&targets,
hierarchy,
&device,
)
.unwrap();
let adamw_params = candle_nn::ParamsAdamW {
lr: 0.01,
..Default::default()
};
let mut optimizer = candle_nn::AdamW::new(varmap.all_vars(), adamw_params).unwrap();
optimizer.backward_step(&loss).unwrap();
let updated_values: Vec<Vec<f32>> = vars
.iter()
.map(|v| {
v.as_tensor()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap()
})
.collect();
let n_changed: usize = initial_values
.iter()
.zip(updated_values.iter())
.filter(|(a, b)| a != b)
.count();
assert!(
n_changed > 5,
"Expected many parameters to change after hierarchical backward, only {} changed out of {}",
n_changed,
vars.len()
);
}
fn hier_training_fixture() -> (Vec<String>, MultiBranchDataset, MultiBranchDataset) {
let labels = make_hier_labels();
let n_classes = labels.len();
let mut records = Vec::new();
for i in 0..30 {
let class_idx = i % n_classes;
let bias = class_idx as f32;
records.push(TrainingRecord {
label: labels[class_idx].clone(),
char_features: (0..960)
.map(|j| {
if j % n_classes == class_idx {
1.0 + bias
} else {
0.1
}
})
.collect(),
embed_features: (0..512)
.map(|j| {
if j % n_classes == class_idx {
1.0 + bias
} else {
0.1
}
})
.collect(),
stats_features: (0..27)
.map(|j| {
if j % n_classes == class_idx {
1.0 + bias
} else {
0.1
}
})
.collect(),
header_features: (0..128)
.map(|j| {
if j % n_classes == class_idx {
0.5 + bias
} else {
0.05
}
})
.collect(),
validation_features: Vec::new(),
values: Vec::new(),
});
}
let mut label_to_idx = std::collections::HashMap::new();
for (i, label) in labels.iter().enumerate() {
label_to_idx.insert(label.clone(), i as u32);
}
let train_data =
MultiBranchDataset::from_records(&records[..20], &label_to_idx, 960, 512, 27, 128)
.unwrap();
let val_data =
MultiBranchDataset::from_records(&records[20..], &label_to_idx, 960, 512, 27, 128)
.unwrap();
(labels, train_data, val_data)
}
fn hier_train_config(
output_dir: &std::path::Path,
epochs: usize,
lr: f64,
min_lr: f64,
) -> MultiBranchTrainConfig {
MultiBranchTrainConfig {
output_dir: output_dir.to_path_buf(),
epochs,
batch_size: 10,
lr,
weight_decay: 1e-4,
patience: 10,
seed: 42,
min_lr,
logit_adjust_tau: 0.0,
}
}
const HIER_VAL_LOSS_CEILING: f32 = 0.6;
#[test]
fn test_hierarchical_training_loop_small() {
let (labels, train_data, val_data) = hier_training_fixture();
let config = MultiBranchConfig {
n_classes: labels.len(),
head_type: HeadType::Hierarchical,
..Default::default()
};
let tmp_dir = tempfile::tempdir().unwrap();
let train_config = hier_train_config(tmp_dir.path(), 5, 1e-3, 1e-6);
let summary = train_multi_branch(
&train_config,
&config,
&train_data,
&val_data,
Some(&labels),
None,
None,
)
.unwrap();
assert_eq!(summary.total_epochs, 5);
let val_first = summary.epoch_metrics[0].val_loss;
let val_last = summary.epoch_metrics[4].val_loss;
assert!(
val_last <= HIER_VAL_LOSS_CEILING * val_first,
"five epochs should cut held-out loss to at most {}x its first epoch: \
epoch 0 = {val_first}, epoch 4 = {val_last} (ratio {:.4}); \
val losses {:?}",
HIER_VAL_LOSS_CEILING,
val_last / val_first,
summary
.epoch_metrics
.iter()
.map(|m| m.val_loss)
.collect::<Vec<_>>(),
);
assert!(tmp_dir.path().join("model.safetensors").exists());
assert!(tmp_dir.path().join("config.json").exists());
}
#[test]
fn test_hierarchical_training_is_reproducible_at_one_seed() {
let (labels, train_data, val_data) = hier_training_fixture();
let config = MultiBranchConfig {
n_classes: labels.len(),
head_type: HeadType::Hierarchical,
..Default::default()
};
let run = || {
let tmp_dir = tempfile::tempdir().unwrap();
let train_config = hier_train_config(tmp_dir.path(), 2, 1e-3, 1e-6);
let summary = train_multi_branch(
&train_config,
&config,
&train_data,
&val_data,
Some(&labels),
None,
None,
)
.unwrap();
summary
.epoch_metrics
.iter()
.map(|m| (m.train_loss, m.val_loss, m.val_accuracy))
.collect::<Vec<_>>()
};
let first = run();
let second = run();
assert_eq!(
first, second,
"same seed, same machine, different numbers: {first:?} then {second:?}"
);
assert!(
first.iter().any(|(train, _, _)| *train > 0.0),
"no train loss was recorded: {first:?}"
);
}
#[test]
fn test_hierarchical_training_criterion_needs_a_working_optimizer() {
let (labels, train_data, val_data) = hier_training_fixture();
let config = MultiBranchConfig {
n_classes: labels.len(),
head_type: HeadType::Hierarchical,
..Default::default()
};
let tmp_dir = tempfile::tempdir().unwrap();
let train_config = hier_train_config(tmp_dir.path(), 2, 0.0, 0.0);
let summary = train_multi_branch(
&train_config,
&config,
&train_data,
&val_data,
Some(&labels),
None,
None,
)
.unwrap();
let val_first = summary.epoch_metrics[0].val_loss;
let val_last = summary.epoch_metrics.last().unwrap().val_loss;
assert!(
val_last > HIER_VAL_LOSS_CEILING * val_first,
"a zero learning rate cut held-out loss from {val_first} to {val_last}, \
which the {}x ceiling would accept — the ceiling is not measuring training",
HIER_VAL_LOSS_CEILING,
);
assert!(
val_last < val_first,
"held-out loss did not fall on a frozen model ({val_first} then {val_last}), \
so `last < first` would have been evidence of training after all and the \
margin above is answering a question nobody is asking",
);
}
#[test]
fn test_ftmb_v3_roundtrip() {
let records: Vec<TrainingRecord> = (0..5)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32 * 1.0; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let groups = vec![
TableGroup {
record_indices: vec![0, 1, 2],
sibling_headers: vec!["city".to_string(), "name".to_string(), "email".to_string()],
},
TableGroup {
record_indices: vec![3, 4],
sibling_headers: vec!["amount".to_string(), "currency".to_string()],
},
];
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data_v3(tmp.path(), &records, &groups, 8, 6, 4, 5).unwrap();
let (header, loaded, loaded_groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 3);
assert_eq!(header.n_records, 5);
assert_eq!(header.n_groups, 2);
assert_eq!(header.char_dim, 8);
assert_eq!(header.embed_dim, 6);
assert_eq!(header.stats_dim, 4);
assert_eq!(header.header_dim, 5);
assert_eq!(loaded.len(), 5);
assert_eq!(loaded_groups.len(), 2);
assert_eq!(loaded_groups[0].record_indices, vec![0, 1, 2]);
assert_eq!(loaded_groups[0].sibling_headers.len(), 3);
assert_eq!(loaded_groups[0].sibling_headers[0], "city");
assert_eq!(loaded_groups[0].sibling_headers[1], "name");
assert_eq!(loaded_groups[0].sibling_headers[2], "email");
assert_eq!(loaded_groups[1].record_indices, vec![3, 4]);
assert_eq!(loaded_groups[1].sibling_headers.len(), 2);
assert_eq!(loaded_groups[1].sibling_headers[0], "amount");
assert_eq!(loaded_groups[1].sibling_headers[1], "currency");
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.embed_features, read.embed_features);
assert_eq!(orig.stats_features, read.stats_features);
assert_eq!(orig.header_features, read.header_features);
}
}
#[test]
fn test_ftmb_v1_read_as_single_group() {
use std::io::Write;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mut file = std::fs::File::create(tmp.path()).unwrap();
let char_dim: u16 = 4;
let embed_dim: u16 = 3;
let stats_dim: u16 = 2;
let n_records: u64 = 2;
file.write_all(b"FTMB").unwrap();
file.write_all(&1u32.to_le_bytes()).unwrap();
file.write_all(&n_records.to_le_bytes()).unwrap();
file.write_all(&char_dim.to_le_bytes()).unwrap();
file.write_all(&embed_dim.to_le_bytes()).unwrap();
file.write_all(&stats_dim.to_le_bytes()).unwrap();
file.write_all(&[0u8; 2]).unwrap();
for i in 0..2 {
let label = format!("type_{i}");
let label_bytes = label.as_bytes();
file.write_all(&(label_bytes.len() as u16).to_le_bytes())
.unwrap();
file.write_all(label_bytes).unwrap();
for j in 0..char_dim {
file.write_all(&((i * 10 + j) as f32).to_le_bytes())
.unwrap();
}
for j in 0..embed_dim {
file.write_all(&((i * 10 + j) as f32 * 0.5).to_le_bytes())
.unwrap();
}
for j in 0..stats_dim {
file.write_all(&((i * 10 + j) as f32 * 0.1).to_le_bytes())
.unwrap();
}
}
drop(file);
let (_header, records, groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(records.len(), 2);
assert_eq!(groups.len(), 1, "v1 should have single group");
assert_eq!(groups[0].record_indices, vec![0, 1]);
assert!(
groups[0].sibling_headers.is_empty(),
"v1 group should have empty sibling headers"
);
}
#[test]
fn test_ftmb_v2_read_as_single_group() {
let records: Vec<TrainingRecord> = (0..3)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data(tmp.path(), &records, 8, 6, 4, 5).unwrap();
let (_header, loaded, groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(loaded.len(), 3);
assert_eq!(groups.len(), 1, "v2 should have single group");
assert_eq!(groups[0].record_indices, vec![0, 1, 2]);
assert!(
groups[0].sibling_headers.is_empty(),
"v2 group should have empty sibling headers"
);
}
#[test]
fn test_frozen_sibling_context_forward() {
use crate::sibling_context::SiblingContextTrainable;
use finetype_model::sibling_context::SiblingContextConfig;
let varmap = VarMap::new();
let config = SiblingContextConfig::default();
let device = Device::Cpu;
let trainable = SiblingContextTrainable::new(&varmap, &config, &device).unwrap();
let tmp_dir = std::env::temp_dir().join("finetype_frozen_sibling_test");
let _ = std::fs::remove_dir_all(&tmp_dir);
std::fs::create_dir_all(&tmp_dir).unwrap();
varmap.save(tmp_dir.join("model.safetensors")).unwrap();
let config_json = serde_json::to_string_pretty(&config).unwrap();
std::fs::write(tmp_dir.join("config.json"), &config_json).unwrap();
let frozen = FrozenSiblingContext::load(&tmp_dir, &device).unwrap();
assert_eq!(frozen.embed_dim(), 128);
for n_cols in [1, 5, 10, 20] {
let input = Tensor::randn(0.0f32, 1.0, (n_cols, 128), &device).unwrap();
let output = frozen.forward(&input).unwrap();
assert_eq!(
output.dims(),
&[n_cols, 128],
"Shape mismatch for N={}",
n_cols
);
}
let input = Tensor::randn(0.0f32, 1.0, (5, 128), &device).unwrap();
let out_trainable: Vec<f32> = trainable
.forward(&input)
.unwrap()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
let out_frozen: Vec<f32> = frozen
.forward(&input)
.unwrap()
.flatten_all()
.unwrap()
.to_vec1()
.unwrap();
for (a, b) in out_trainable.iter().zip(out_frozen.iter()) {
assert!(
(a - b).abs() < 1e-5,
"Frozen output should match trainable: {} vs {}",
a,
b
);
}
let _ = std::fs::remove_dir_all(&tmp_dir);
}
#[test]
fn test_dataset_batch_groups_without_enrichment() {
let records: Vec<TrainingRecord> = (0..6)
.map(|i| TrainingRecord {
label: format!("identity.person.type_{}", i % 3),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let mut label_to_idx = std::collections::HashMap::new();
label_to_idx.insert("identity.person.type_0".to_string(), 0u32);
label_to_idx.insert("identity.person.type_1".to_string(), 1u32);
label_to_idx.insert("identity.person.type_2".to_string(), 2u32);
let groups = vec![
TableGroup {
record_indices: vec![0, 1, 2],
sibling_headers: vec!["city".to_string(), "name".to_string(), "email".to_string()],
},
TableGroup {
record_indices: vec![3, 4, 5],
sibling_headers: vec![
"amount".to_string(),
"currency".to_string(),
"date".to_string(),
],
},
];
let dataset = MultiBranchDataset::from_records_with_groups(
&records,
&label_to_idx,
8,
6,
4,
5,
0,
Some(groups),
)
.unwrap();
let device = Device::Cpu;
let (char_t, embed_t, stats_t, header_t, _valid_t, labels_t) =
dataset.batch_groups(&[0], None, &device).unwrap();
assert_eq!(char_t.dims(), &[3, 8]);
assert_eq!(embed_t.dims(), &[3, 6]);
assert_eq!(stats_t.dims(), &[3, 4]);
let header_t = header_t.expect("header should be Some");
assert_eq!(header_t.dims(), &[3, 5]);
assert_eq!(labels_t.dims(), &[3]);
let (char_t, _, _, _, _, labels_t) = dataset.batch_groups(&[0, 1], None, &device).unwrap();
assert_eq!(char_t.dims(), &[6, 8]);
assert_eq!(labels_t.dims(), &[6]);
}
#[test]
fn test_ftmb_v4_roundtrip() {
let valid_dim: usize = 239;
let records: Vec<TrainingRecord> = (0..10)
.map(|i| {
let validation_features: Vec<f32> = (0..valid_dim)
.map(|j| if (i + j) % 2 == 0 { 0.0 } else { 1.0 })
.collect();
TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32 * 1.0; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features,
values: Vec::new(),
}
})
.collect();
let groups = vec![
TableGroup {
record_indices: vec![0, 1, 2, 3, 4],
sibling_headers: vec![
"a".to_string(),
"b".to_string(),
"c".to_string(),
"d".to_string(),
"e".to_string(),
],
},
TableGroup {
record_indices: vec![5, 6, 7, 8, 9],
sibling_headers: vec![
"f".to_string(),
"g".to_string(),
"h".to_string(),
"i".to_string(),
"j".to_string(),
],
},
];
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data_v4(tmp.path(), &records, &groups, 8, 6, 4, 5, valid_dim as u16)
.unwrap();
let (header, loaded, loaded_groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 4);
assert_eq!(header.n_records, 10);
assert_eq!(header.char_dim, 8);
assert_eq!(header.embed_dim, 6);
assert_eq!(header.stats_dim, 4);
assert_eq!(header.header_dim, 5);
assert_eq!(header.n_groups, 2);
assert_eq!(header.valid_dim, valid_dim as u16);
assert_eq!(loaded.len(), 10);
assert_eq!(loaded_groups.len(), 2);
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.embed_features, read.embed_features);
assert_eq!(orig.stats_features, read.stats_features);
assert_eq!(orig.header_features, read.header_features);
assert_eq!(
orig.validation_features, read.validation_features,
"Validation features must round-trip exactly for record {}",
orig.label
);
assert_eq!(read.validation_features.len(), valid_dim);
}
assert_eq!(loaded_groups[0].record_indices, vec![0, 1, 2, 3, 4]);
assert_eq!(loaded_groups[1].record_indices, vec![5, 6, 7, 8, 9]);
}
#[test]
fn test_ftmb_v6_roundtrip_values() {
let valid_dim: usize = 12;
let n_values: u16 = 4;
let value_sets: Vec<Vec<String>> = vec![
vec!["user@x.com".into(), "bob@y.org".into()],
vec!["£42.50".into(), "naïve".into(), "日本語".into(), "x".into()],
vec![], vec![
"a".into(),
"b".into(),
"c".into(),
"d".into(),
"e".into(),
"f".into(),
], ];
let records: Vec<TrainingRecord> = (0..4)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: vec![if i % 2 == 0 { 0.0 } else { 1.0 }; valid_dim],
values: value_sets[i].clone(),
})
.collect();
let groups = vec![TableGroup {
record_indices: vec![0, 1, 2, 3],
sibling_headers: vec!["a".into(), "b".into(), "c".into(), "d".into()],
}];
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data_v6(
tmp.path(),
&records,
&groups,
8,
6,
4,
5,
valid_dim as u16,
n_values,
)
.unwrap();
let raw = std::fs::read(tmp.path()).unwrap();
assert_eq!(&raw[0..4], b"FTMB");
assert_eq!(u32::from_le_bytes([raw[4], raw[5], raw[6], raw[7]]), 6);
let (header, loaded, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 6);
assert_eq!(loaded.len(), 4);
for (i, (orig, read)) in records.iter().zip(loaded.iter()).enumerate() {
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.validation_features, read.validation_features);
let expected: Vec<String> = orig
.values
.iter()
.take(n_values as usize)
.cloned()
.collect();
assert_eq!(read.values, expected, "values mismatch for record {i}");
}
assert_eq!(loaded[1].values, vec!["£42.50", "naïve", "日本語", "x"]);
assert_eq!(loaded[2].values, Vec::<String>::new());
assert_eq!(loaded[3].values, vec!["a", "b", "c", "d"]);
}
#[test]
fn test_ftmb_v4_header_is_30_bytes() {
assert_eq!(FTMB_HEADER_SIZE_V4, 30);
let records = vec![TrainingRecord {
label: "test.type.t_0".to_string(),
char_features: vec![1.0; 4],
embed_features: vec![2.0; 3],
stats_features: vec![3.0; 2],
header_features: vec![4.0; 5],
validation_features: vec![0.5; 10],
values: Vec::new(),
}];
let groups = vec![TableGroup {
record_indices: vec![0],
sibling_headers: vec!["col".to_string()],
}];
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data_v4(tmp.path(), &records, &groups, 4, 3, 2, 5, 10).unwrap();
let raw = std::fs::read(tmp.path()).unwrap();
assert_eq!(&raw[0..4], b"FTMB");
assert_eq!(u32::from_le_bytes([raw[4], raw[5], raw[6], raw[7]]), 4);
assert_eq!(
u16::from_le_bytes([raw[28], raw[29]]),
10,
"valid_dim should be at offset 28"
);
}
#[test]
fn test_ftmb_v4_reader_v1_compat_validation_zeros() {
use std::io::Write;
let tmp = tempfile::NamedTempFile::new().unwrap();
let mut file = std::fs::File::create(tmp.path()).unwrap();
let char_dim: u16 = 4;
let embed_dim: u16 = 3;
let stats_dim: u16 = 2;
let n_records: u64 = 2;
file.write_all(b"FTMB").unwrap();
file.write_all(&1u32.to_le_bytes()).unwrap();
file.write_all(&n_records.to_le_bytes()).unwrap();
file.write_all(&char_dim.to_le_bytes()).unwrap();
file.write_all(&embed_dim.to_le_bytes()).unwrap();
file.write_all(&stats_dim.to_le_bytes()).unwrap();
file.write_all(&[0u8; 2]).unwrap();
for i in 0..2u16 {
let label = format!("type_{i}");
let label_bytes = label.as_bytes();
file.write_all(&(label_bytes.len() as u16).to_le_bytes())
.unwrap();
file.write_all(label_bytes).unwrap();
for j in 0..char_dim {
file.write_all(&((i * 10 + j) as f32).to_le_bytes())
.unwrap();
}
for j in 0..embed_dim {
file.write_all(&((i * 10 + j) as f32 * 0.5).to_le_bytes())
.unwrap();
}
for j in 0..stats_dim {
file.write_all(&((i * 10 + j) as f32 * 0.1).to_le_bytes())
.unwrap();
}
}
drop(file);
let (header, records, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 1);
assert_eq!(header.valid_dim, 0);
for record in &records {
assert!(
record.validation_features.is_empty(),
"v1 files should produce empty validation_features"
);
assert!(
record.header_features.is_empty(),
"v1 files should produce empty header_features"
);
}
}
#[test]
fn test_ftmb_v4_reader_v2_compat_validation_zeros() {
let records: Vec<TrainingRecord> = (0..3)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data(tmp.path(), &records, 8, 6, 4, 5).unwrap();
let (header, loaded, _groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 2);
assert_eq!(header.valid_dim, 0);
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.embed_features, read.embed_features);
assert_eq!(orig.stats_features, read.stats_features);
assert_eq!(orig.header_features, read.header_features);
assert!(
read.validation_features.is_empty(),
"v2 files should produce empty validation_features"
);
}
}
#[test]
fn test_ftmb_v4_reader_v3_compat_validation_zeros() {
let records: Vec<TrainingRecord> = (0..4)
.map(|i| TrainingRecord {
label: format!("test.type.t_{i}"),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: Vec::new(),
values: Vec::new(),
})
.collect();
let groups = vec![
TableGroup {
record_indices: vec![0, 1],
sibling_headers: vec!["x".to_string(), "y".to_string()],
},
TableGroup {
record_indices: vec![2, 3],
sibling_headers: vec!["a".to_string(), "b".to_string()],
},
];
let tmp = tempfile::NamedTempFile::new().unwrap();
write_training_data_v3(tmp.path(), &records, &groups, 8, 6, 4, 5).unwrap();
let (header, loaded, loaded_groups) = read_training_data(tmp.path()).unwrap();
assert_eq!(header.version, 3);
assert_eq!(header.valid_dim, 0);
assert_eq!(loaded_groups.len(), 2);
for (orig, read) in records.iter().zip(loaded.iter()) {
assert_eq!(orig.label, read.label);
assert_eq!(orig.char_features, read.char_features);
assert_eq!(orig.embed_features, read.embed_features);
assert_eq!(orig.stats_features, read.stats_features);
assert_eq!(orig.header_features, read.header_features);
assert!(
read.validation_features.is_empty(),
"v3 files should produce empty validation_features"
);
}
}
#[test]
fn test_forward_pass_with_validation_branch() {
let config = MultiBranchConfig {
valid_dim: 239,
valid_hidden: [128, 64],
..MultiBranchConfig::default()
};
assert!(config.has_validation_branch());
assert_eq!(config.merged_dim(), 692);
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 5;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let valid_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.valid_dim), &device).unwrap();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
Some(&valid_feats),
false,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
Some(&valid_feats),
true,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
}
#[test]
fn test_forward_pass_without_validation_branch() {
let config = MultiBranchConfig::default();
assert!(!config.has_validation_branch());
assert_eq!(config.valid_dim, 0);
assert_eq!(config.merged_dim(), 628);
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 5;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
None,
false,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
}
#[test]
fn test_dataset_batch_with_validation_features() {
let valid_dim = 10;
let records: Vec<TrainingRecord> = (0..5)
.map(|i| TrainingRecord {
label: format!("identity.person.type_{}", i % 3),
char_features: vec![i as f32; 8],
embed_features: vec![i as f32 * 2.0; 6],
stats_features: vec![i as f32 * 3.0; 4],
header_features: vec![i as f32 * 4.0; 5],
validation_features: (0..valid_dim).map(|j| (i + j) as f32 * 0.1).collect(),
values: Vec::new(),
})
.collect();
let mut label_to_idx = std::collections::HashMap::new();
label_to_idx.insert("identity.person.type_0".to_string(), 0u32);
label_to_idx.insert("identity.person.type_1".to_string(), 1u32);
label_to_idx.insert("identity.person.type_2".to_string(), 2u32);
let dataset = MultiBranchDataset::from_records_with_groups(
&records,
&label_to_idx,
8,
6,
4,
5,
valid_dim,
None,
)
.unwrap();
let device = Device::Cpu;
let (char_t, embed_t, stats_t, header_t, valid_t, labels_t) =
dataset.batch(&[0, 1, 2, 3, 4], &device).unwrap();
assert_eq!(char_t.dims(), &[5, 8]);
assert_eq!(embed_t.dims(), &[5, 6]);
assert_eq!(stats_t.dims(), &[5, 4]);
let header_t = header_t.expect("header should be Some");
assert_eq!(header_t.dims(), &[5, 5]);
let valid_t = valid_t.expect("valid should be Some when valid_dim > 0");
assert_eq!(valid_t.dims(), &[5, valid_dim]);
assert_eq!(labels_t.dims(), &[5]);
}
#[test]
fn test_v13_validation_branch_resize() {
let config = MultiBranchConfig {
char_hidden: [450, 450],
embed_hidden: [300, 300],
stats_hidden: [192, 96],
header_hidden: [192, 96],
valid_dim: 239,
valid_hidden: [192, 128],
merge_hidden: [750, 750],
..MultiBranchConfig::default()
};
assert!(config.has_validation_branch());
assert_eq!(config.merged_dim(), 1070);
let device = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 4;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let valid_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.valid_dim), &device).unwrap();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
Some(&valid_feats),
false,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
Some(&valid_feats),
true,
)
.unwrap();
assert_eq!(logits.dims(), &[batch_size, config.n_classes]);
}
#[test]
fn test_v13_config_json_deserialization() {
let json = r#"{
"char_dim": 960,
"embed_dim": 512,
"stats_dim": 27,
"header_dim": 128,
"valid_dim": 240,
"char_hidden": [450, 450],
"embed_hidden": [300, 300],
"stats_hidden": [192, 96],
"header_hidden": [192, 96],
"valid_hidden": [192, 128],
"merge_hidden": [750, 750],
"n_classes": 240,
"dropout": 0.35,
"head_type": "Flat",
"activation": "ReLU",
"use_layer_norm": false
}"#;
let config: MultiBranchConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.valid_dim, 240);
assert_eq!(config.valid_hidden, [192, 128]);
assert!(config.has_validation_branch());
assert_eq!(config.merged_dim(), 1070);
}
#[test]
fn test_default_config_valid_hidden_backward_compat() {
let config = MultiBranchConfig::default();
assert_eq!(config.valid_hidden, [0, 0]);
assert_eq!(config.valid_dim, 0);
assert!(!config.has_validation_branch());
}
#[test]
fn test_gradient_norm_computation() {
let device = Device::Cpu;
let config = MultiBranchConfig {
valid_dim: 239,
valid_hidden: [128, 64],
..MultiBranchConfig::default()
};
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &device);
let model = MultiBranchModel::new(&config, vb).unwrap();
let batch_size = 4;
let char_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.char_dim), &device).unwrap();
let embed_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.embed_dim), &device).unwrap();
let stats_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.stats_dim), &device).unwrap();
let header_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.header_dim), &device).unwrap();
let valid_feats =
Tensor::randn(0.0f32, 1.0, (batch_size, config.valid_dim), &device).unwrap();
let logits = model
.forward(
&char_feats,
&embed_feats,
&stats_feats,
Some(&header_feats),
Some(&valid_feats),
true,
)
.unwrap();
let targets = Tensor::new(&[0u32, 1, 2, 3], &device).unwrap();
let loss = candle_nn::loss::cross_entropy(&logits, &targets).unwrap();
let grads = loss.backward().unwrap();
let norms = compute_branch_gradient_norms(&varmap, &grads);
assert!(norms.contains_key("char"), "char branch missing from norms");
assert!(
norms.contains_key("embed"),
"embed branch missing from norms"
);
assert!(
norms.contains_key("stats"),
"stats branch missing from norms"
);
assert!(
norms.contains_key("header"),
"header branch missing from norms"
);
assert!(
norms.contains_key("valid"),
"valid branch missing from norms"
);
assert!(
norms.contains_key("merge"),
"merge branch missing from norms"
);
assert!(norms.contains_key("head"), "head branch missing from norms");
for (branch, norm) in &norms {
assert!(
*norm > 0.0 && norm.is_finite(),
"branch {} has invalid norm: {}",
branch,
norm
);
}
}
}