Skip to main content

zsfm_tabfm/infer/
mod.rs

1//! TabFM inference engine, with a real batch dimension `B` (`predict_batch`) — e.g. one batch
2//! item per ensemble member, so `ensemble::orchestrate` can share the fixed cost of the 24-block
3//! ICL stage across all members in one forward pass instead of paying it once per member. The
4//! single-table `predict()` is a thin `B=1` wrapper around the same code path; none of the
5//! attention/RMSNorm/RoPE math below changed to add batching — only the four "stage" functions
6//! (`cell_embed`, `col_embedding_forward`, `row_interaction_forward`, `icl_forward`) gained a
7//! leading `B` dimension, via reshapes around the *same* 3D attention calls they always made
8//! (masks/weights are shared scalars across the batch — every member has the same row/feature
9//! count, `train_size`, and `d`; only cell *values* and `cat_mask` vary per member).
10//!
11//! Architecture (from `tabfm/src/pytorch/model.py`, verified against the installed package):
12//! CellEmbedder (per-cell grouped Fourier features + train-row y-embedding)
13//!   -> ColEmbedding (SetTransformer / induced attention over rows, masked to train rows)
14//!   -> prepend row_num_cls learned CLS tokens on the column axis
15//!   -> RowInteraction (RoPE cross-column self-attention, masked to valid/unpadded columns)
16//!   -> ColEmbedding (stage 2)
17//!   -> RowInteraction (stage 2, output collapsed to the CLS-token slice -> icl_dim)
18//!   -> ICLearning (24-block self-attention over rows, y re-injected at train rows, masked so
19//!      only train rows are attendable keys) -> MLP decoder -> per-class logits or a scalar.
20//!
21//! Numeric details that matter for parity (see model-to-gguf skill's debugging ladder):
22//! * RoPE frequencies are checkpoint-loaded buffers, never recomputed from a formula.
23//! * `MultiheadAttention` pre-scales `q` by a learned per-dimension softplus'd scale, then calls
24//!   attention with `scale=1.0` — do not apply an additional `1/sqrt(d)`.
25//! * All masking is additive key-side masking (no causal masking anywhere).
26//! * RoPE here is the *interleaved-pair* variant (`x[0::2]`, `x[1::2]`), NOT the Llama
27//!   rotate-half variant used elsewhere in this repo (see `toto`'s `infer/rope.rs`) —
28//!   deliberately not reused.
29
30use std::io::{Read, Seek};
31use std::path::{Path, PathBuf};
32
33use anyhow::{Context, Result};
34use candle_core::quantized::gguf_file;
35use candle_core::{DType, Device, Tensor, D};
36use zsfm_nn::{linear, load_weight};
37
38use crate::config::TabFMConfig;
39
40// ---------------------------------------------------------------------------
41// Config
42// ---------------------------------------------------------------------------
43
44#[derive(Clone, Debug)]
45pub struct InferConfig {
46    embed_dim: usize,
47    max_classes: usize,
48    col_num_blocks: usize,
49    col_nhead: usize,
50    /// Kept for config-shape parity with `config.json`; the SetTransformer's inducing-point
51    /// count is implied by the loaded `ind_vectors` tensor's own shape, not read from here.
52    #[allow(dead_code)]
53    col_num_inds: usize,
54    row_num_blocks: usize,
55    row_nhead: usize,
56    row_num_cls: usize,
57    icl_num_blocks: usize,
58    icl_nhead: usize,
59    ff_factor: usize,
60    feature_group_size: usize,
61    num_freq: usize,
62    norm_eps: f64,
63    is_classifier: bool,
64    /// `None` means "use the `TabFM.__init__` default of `icl_dim * 2`".
65    decoder_hidden: Option<usize>,
66}
67
68/// Every field here is load-bearing architecture metadata (must match the checkpoint the GGUF
69/// was converted from) — there is no sensible standalone `Default`. Build one from a parsed
70/// `config.json` via `From<&TabFMConfig>` (or `TabFMModelBuilder::config_from`), then optionally
71/// layer on the few genuinely optional overrides below.
72impl From<&TabFMConfig> for InferConfig {
73    fn from(tc: &TabFMConfig) -> Self {
74        InferConfig {
75            embed_dim: tc.embed_dim as usize,
76            max_classes: tc.max_classes as usize,
77            col_num_blocks: tc.col_num_blocks as usize,
78            col_nhead: tc.col_nhead as usize,
79            col_num_inds: tc.col_num_inds as usize,
80            row_num_blocks: tc.row_num_blocks as usize,
81            row_nhead: tc.row_nhead as usize,
82            row_num_cls: tc.row_num_cls as usize,
83            icl_num_blocks: tc.icl_num_blocks as usize,
84            icl_nhead: tc.icl_nhead as usize,
85            ff_factor: tc.ff_factor as usize,
86            feature_group_size: tc.feature_group_size as usize,
87            num_freq: tc.num_freq as usize,
88            norm_eps: tc.norm_eps,
89            is_classifier: tc.is_classifier,
90            decoder_hidden: tc.decoder_hidden.map(|v| v as usize),
91        }
92    }
93}
94
95impl InferConfig {
96    fn col_dim_ff(&self) -> usize { self.embed_dim * self.ff_factor }
97    fn icl_dim(&self) -> usize { self.embed_dim * self.row_num_cls }
98    fn icl_dim_ff(&self) -> usize { self.icl_dim() * self.ff_factor }
99    fn decoder_hidden(&self) -> usize { self.decoder_hidden.unwrap_or(self.icl_dim() * 2) }
100    fn out_dim(&self) -> usize { if self.is_classifier { self.max_classes } else { 1 } }
101
102    // -- builder-style overrides of the few genuinely optional fields ---------
103
104    pub fn with_decoder_hidden(mut self, v: Option<usize>) -> Self { self.decoder_hidden = v; self }
105    pub fn with_norm_eps(mut self, v: f64) -> Self { self.norm_eps = v; self }
106    pub fn with_num_freq(mut self, v: usize) -> Self { self.num_freq = v; self }
107}
108
109// ---------------------------------------------------------------------------
110// Builder
111// ---------------------------------------------------------------------------
112
113/// Fluent constructor for [`TabFMModel`]: point it at a GGUF file and a config (from a parsed
114/// `config.json` via [`config_from`](TabFMModelBuilder::config_from), or a hand-built
115/// [`InferConfig`] via [`config`](TabFMModelBuilder::config)), then call
116/// [`build`](TabFMModelBuilder::build).
117///
118/// ```no_run
119/// use zsfm_tabfm::{TabFMConfig, TabFMModel};
120///
121/// # fn main() -> anyhow::Result<()> {
122/// let tc = TabFMConfig::from_json(&std::fs::read_to_string("config.json")?)?;
123/// let model = TabFMModel::builder("tabfm.gguf").config_from(&tc).build()?;
124/// # Ok(()) }
125/// ```
126pub struct TabFMModelBuilder {
127    gguf_path: PathBuf,
128    config: Option<InferConfig>,
129}
130
131impl TabFMModelBuilder {
132    fn new(gguf_path: impl Into<PathBuf>) -> Self {
133        Self { gguf_path: gguf_path.into(), config: None }
134    }
135
136    /// Use an already-built [`InferConfig`] (e.g. `InferConfig::from(&tc).with_norm_eps(...)`).
137    pub fn config(mut self, config: InferConfig) -> Self {
138        self.config = Some(config);
139        self
140    }
141
142    /// Map a parsed `config.json` (`TabFMConfig`) onto the model's `InferConfig`.
143    pub fn config_from(mut self, tc: &TabFMConfig) -> Self {
144        self.config = Some(InferConfig::from(tc));
145        self
146    }
147
148    pub fn build(self) -> Result<TabFMModel> {
149        let config = self
150            .config
151            .context("TabFMModelBuilder: no config set — call .config(...) or .config_from(...)")?;
152        TabFMModel::load(&self.gguf_path, config)
153    }
154}
155
156// ---------------------------------------------------------------------------
157// Weight structs
158// ---------------------------------------------------------------------------
159
160/// Weights for one `MultiheadAttentionBlock` (attention sublayer + SwiGLU FFN sublayer, each
161/// with its own pre/post RMSNorm).
162struct MabWeights {
163    q_w: Tensor, q_b: Tensor,
164    k_w: Tensor, k_b: Tensor,
165    v_w: Tensor, v_b: Tensor,
166    o_w: Tensor, o_b: Tensor,
167    q_norm: Tensor,
168    k_norm: Tensor,
169    per_dim_scale: Tensor,
170    pre_attn_norm: Tensor,
171    post_attn_norm: Tensor,
172    pre_ff_norm: Tensor,
173    post_ff_norm: Tensor,
174    ffn_up_w: Tensor, ffn_up_b: Tensor,
175    ffn_gate_w: Tensor, ffn_gate_b: Tensor,
176    ffn_down_w: Tensor, ffn_down_b: Tensor,
177}
178
179/// One `InducedSelfAttentionBlock`: shared induced vectors + two `MultiheadAttentionBlock`s.
180struct InducedBlockWeights {
181    ind_vectors: Tensor, // [num_inds, E]
182    mab1: MabWeights,
183    mab2: MabWeights,
184}
185
186struct ColStackWeights {
187    blocks: Vec<InducedBlockWeights>,
188    out_w_w: Tensor, out_w_b: Tensor,
189    out_norm_w: Tensor,
190}
191
192struct RowStackWeights {
193    rope_freqs: Tensor, // [head_dim/2]
194    blocks: Vec<MabWeights>,
195    out_norm_w: Tensor,
196}
197
198/// A generic (weight, bias) linear-layer stack, GELU-tanh between layers (not after the last).
199struct MlpWeights {
200    layers: Vec<(Tensor, Tensor)>,
201}
202
203enum YEmbedWeights {
204    /// classification: plain `nn.Embedding` lookup table `[max_classes, E]`.
205    Embedding(Tensor),
206    /// regression: 2-layer MLP `1 -> 6 -> E`.
207    Mlp(MlpWeights),
208}
209
210enum YEncoderWeights {
211    /// classification: `OneHotAndLinear` projection `[icl_dim, max_classes]` + bias.
212    OneHot { proj_w: Tensor, proj_b: Tensor },
213    /// regression: 2-layer MLP `1 -> decoder_hidden -> icl_dim`.
214    Mlp(MlpWeights),
215}
216
217struct CellWeights {
218    fourier_freq: Tensor,     // [feature_group_size, num_freq]
219    fourier_freq_cat: Tensor, // [feature_group_size, num_freq]
220    in_linear_w: Tensor, in_linear_b: Tensor,         // [E, 2*num_freq], [E]
221    in_linear_cat_w: Tensor, in_linear_cat_b: Tensor, // [E, 2*num_freq], [E]
222    y_embed: YEmbedWeights,
223}
224
225struct IclWeights {
226    blocks: Vec<MabWeights>,
227    out_norm_w: Tensor,
228    y_encoder: YEncoderWeights,
229    decoder: MlpWeights,
230}
231
232pub struct TabFMModel {
233    device: Device,
234    config: InferConfig,
235    cell: CellWeights,
236    colenc1: ColStackWeights,
237    colenc2: ColStackWeights,
238    rowenc1: RowStackWeights,
239    rowenc2: RowStackWeights,
240    cls_tokens: Tensor, // [row_num_cls, E]
241    icl: IclWeights,
242}
243
244// ---------------------------------------------------------------------------
245// GGUF loading helpers
246// ---------------------------------------------------------------------------
247
248fn load_tensor(
249    content: &gguf_file::Content,
250    reader: &mut (impl Read + Seek),
251    name: &str,
252    device: &Device,
253) -> Result<Tensor> {
254    zsfm_nn::load_tensor(content, reader, name, device, DType::F32)
255}
256
257fn load_mab(
258    content: &gguf_file::Content,
259    reader: &mut (impl Read + Seek),
260    prefix: &str,
261    e: usize,
262    ff: usize,
263    device: &Device,
264) -> Result<MabWeights> {
265    let p = |s: &str| format!("{prefix}.{s}");
266    Ok(MabWeights {
267        q_w: load_weight(content, reader, &p("attn_q.weight"), e, device)?,
268        q_b: load_tensor(content, reader, &p("attn_q.bias"), device)?,
269        k_w: load_weight(content, reader, &p("attn_k.weight"), e, device)?,
270        k_b: load_tensor(content, reader, &p("attn_k.bias"), device)?,
271        v_w: load_weight(content, reader, &p("attn_v.weight"), e, device)?,
272        v_b: load_tensor(content, reader, &p("attn_v.bias"), device)?,
273        o_w: load_weight(content, reader, &p("attn_o.weight"), e, device)?,
274        o_b: load_tensor(content, reader, &p("attn_o.bias"), device)?,
275        q_norm: load_tensor(content, reader, &p("q_norm.weight"), device)?,
276        k_norm: load_tensor(content, reader, &p("k_norm.weight"), device)?,
277        per_dim_scale: load_tensor(content, reader, &p("per_dim_scale"), device)?,
278        pre_attn_norm: load_tensor(content, reader, &p("pre_attn_norm.weight"), device)?,
279        post_attn_norm: load_tensor(content, reader, &p("post_attn_norm.weight"), device)?,
280        pre_ff_norm: load_tensor(content, reader, &p("pre_ff_norm.weight"), device)?,
281        post_ff_norm: load_tensor(content, reader, &p("post_ff_norm.weight"), device)?,
282        ffn_up_w: load_weight(content, reader, &p("ffn_up.weight"), ff, device)?,
283        ffn_up_b: load_tensor(content, reader, &p("ffn_up.bias"), device)?,
284        ffn_gate_w: load_weight(content, reader, &p("ffn_gate.weight"), ff, device)?,
285        ffn_gate_b: load_tensor(content, reader, &p("ffn_gate.bias"), device)?,
286        ffn_down_w: load_weight(content, reader, &p("ffn_down.weight"), e, device)?,
287        ffn_down_b: load_tensor(content, reader, &p("ffn_down.bias"), device)?,
288    })
289}
290
291fn load_col_stack(
292    content: &gguf_file::Content,
293    reader: &mut (impl Read + Seek),
294    stack: &str,
295    num_blocks: usize,
296    e: usize,
297    ff: usize,
298    device: &Device,
299) -> Result<ColStackWeights> {
300    let mut blocks = Vec::with_capacity(num_blocks);
301    for n in 0..num_blocks {
302        let ind_vectors = load_tensor(content, reader, &format!("{stack}.blk.{n}.ind_vectors"), device)?;
303        let mab1 = load_mab(content, reader, &format!("{stack}.blk.{n}.mab1"), e, ff, device)?;
304        let mab2 = load_mab(content, reader, &format!("{stack}.blk.{n}.mab2"), e, ff, device)?;
305        blocks.push(InducedBlockWeights { ind_vectors, mab1, mab2 });
306    }
307    Ok(ColStackWeights {
308        blocks,
309        out_w_w: load_weight(content, reader, &format!("{stack}.out_w.weight"), e, device)?,
310        out_w_b: load_tensor(content, reader, &format!("{stack}.out_w.bias"), device)?,
311        out_norm_w: load_tensor(content, reader, &format!("{stack}.out_norm.weight"), device)?,
312    })
313}
314
315fn load_row_stack(
316    content: &gguf_file::Content,
317    reader: &mut (impl Read + Seek),
318    stack: &str,
319    num_blocks: usize,
320    e: usize,
321    ff: usize,
322    device: &Device,
323) -> Result<RowStackWeights> {
324    let rope_freqs = load_tensor(content, reader, &format!("{stack}.rope_freqs"), device)?;
325    let mut blocks = Vec::with_capacity(num_blocks);
326    for n in 0..num_blocks {
327        blocks.push(load_mab(content, reader, &format!("{stack}.blk.{n}"), e, ff, device)?);
328    }
329    Ok(RowStackWeights {
330        rope_freqs,
331        blocks,
332        out_norm_w: load_tensor(content, reader, &format!("{stack}.out_norm.weight"), device)?,
333    })
334}
335
336fn load_mlp(
337    content: &gguf_file::Content,
338    reader: &mut (impl Read + Seek),
339    prefix: &str,
340    dims: &[usize], // e.g. [in, hidden, out] -> layers [(in,hidden), (hidden,out)]
341    device: &Device,
342) -> Result<MlpWeights> {
343    let mut layers = Vec::with_capacity(dims.len() - 1);
344    for i in 0..dims.len() - 1 {
345        let out_dim = dims[i + 1];
346        let w = load_weight(content, reader, &format!("{prefix}.mlp.{i}.weight"), out_dim, device)?;
347        let b = load_tensor(content, reader, &format!("{prefix}.mlp.{i}.bias"), device)?;
348        layers.push((w, b));
349    }
350    Ok(MlpWeights { layers })
351}
352
353impl TabFMModel {
354    /// Start building a [`TabFMModel`] — see [`TabFMModelBuilder`].
355    pub fn builder(gguf_path: impl Into<PathBuf>) -> TabFMModelBuilder {
356        TabFMModelBuilder::new(gguf_path)
357    }
358
359    pub fn load(gguf_path: &Path, config: InferConfig) -> Result<Self> {
360        let device = Device::Cpu;
361        let mut file = std::fs::File::open(gguf_path)
362            .with_context(|| format!("open {}", gguf_path.display()))?;
363        let content = gguf_file::Content::read(&mut file).context("parse GGUF header")?;
364
365        let e = config.embed_dim;
366        let col_ff = config.col_dim_ff();
367        let icl_dim = config.icl_dim();
368        let icl_ff = config.icl_dim_ff();
369        let decoder_hidden = config.decoder_hidden();
370        let out_dim = config.out_dim();
371
372        let cell = CellWeights {
373            fourier_freq: load_tensor(&content, &mut file, "cell.fourier_freq", &device)?,
374            fourier_freq_cat: load_tensor(&content, &mut file, "cell.fourier_freq_cat", &device)?,
375            in_linear_w: load_weight(&content, &mut file, "cell.in_linear.weight", e, &device)?,
376            in_linear_b: load_tensor(&content, &mut file, "cell.in_linear.bias", &device)?,
377            in_linear_cat_w: load_weight(&content, &mut file, "cell.in_linear_cat.weight", e, &device)?,
378            in_linear_cat_b: load_tensor(&content, &mut file, "cell.in_linear_cat.bias", &device)?,
379            y_embed: if config.is_classifier {
380                YEmbedWeights::Embedding(load_tensor(&content, &mut file, "cell.y_embed.weight", &device)?)
381            } else {
382                YEmbedWeights::Mlp(load_mlp(&content, &mut file, "cell.y_embed", &[1, 6, e], &device)?)
383            },
384        };
385
386        let colenc1 = load_col_stack(&content, &mut file, "colenc1", config.col_num_blocks, e, col_ff, &device)?;
387        let colenc2 = load_col_stack(&content, &mut file, "colenc2", config.col_num_blocks, e, col_ff, &device)?;
388        let rowenc1 = load_row_stack(&content, &mut file, "rowenc1", config.row_num_blocks, e, col_ff, &device)?;
389        let rowenc2 = load_row_stack(&content, &mut file, "rowenc2", config.row_num_blocks, e, col_ff, &device)?;
390
391        let cls_tokens = load_tensor(&content, &mut file, "cls_tokens", &device)?;
392
393        let mut icl_blocks = Vec::with_capacity(config.icl_num_blocks);
394        for n in 0..config.icl_num_blocks {
395            icl_blocks.push(load_mab(&content, &mut file, &format!("icl.blk.{n}"), icl_dim, icl_ff, &device)?);
396        }
397        let y_encoder = if config.is_classifier {
398            YEncoderWeights::OneHot {
399                proj_w: load_weight(&content, &mut file, "icl.y_encoder.projection.weight", icl_dim, &device)?,
400                proj_b: load_tensor(&content, &mut file, "icl.y_encoder.projection.bias", &device)?,
401            }
402        } else {
403            YEncoderWeights::Mlp(load_mlp(&content, &mut file, "icl.y_encoder", &[1, decoder_hidden, icl_dim], &device)?)
404        };
405        let decoder = load_mlp(&content, &mut file, "icl.decoder", &[icl_dim, decoder_hidden, out_dim], &device)?;
406        let icl = IclWeights {
407            blocks: icl_blocks,
408            out_norm_w: load_tensor(&content, &mut file, "icl.out_norm.weight", &device)?,
409            y_encoder,
410            decoder,
411        };
412
413        Ok(Self { device, config, cell, colenc1, colenc2, rowenc1, rowenc2, cls_tokens, icl })
414    }
415
416    // -----------------------------------------------------------------------
417    // Public API
418    // -----------------------------------------------------------------------
419
420    pub fn is_classifier(&self) -> bool {
421        self.config.is_classifier
422    }
423
424    /// Run one table (train rows followed by test rows) through TabFM. A thin `B=1` wrapper
425    /// around `predict_batch` — see that method to run many tables (e.g. ensemble members) in
426    /// one forward pass.
427    ///
428    /// `x`: `[T][H]` padded feature matrix (numeric; categorical columns are pre-encoded to
429    /// floats by the caller). `y`: `[T]` labels (any finite placeholder at test-row positions is
430    /// fine — it's masked internally and never influences the output). `train_size`: number of
431    /// leading rows that are training rows. `cat_mask`: `[H]`, which columns are categorical
432    /// (defaults to all-false). `d`: actual (unpadded) feature count (defaults to `H`).
433    ///
434    /// Returns `[T][out_dim]` raw logits (classification, `out_dim = max_classes`) or a
435    /// `[T][1]` scalar (regression) — only rows `>= train_size` are meaningful predictions.
436    pub fn predict(
437        &self,
438        x: &[Vec<f32>],
439        y: &[f32],
440        train_size: usize,
441        cat_mask: Option<&[bool]>,
442        d: Option<usize>,
443    ) -> Result<Vec<Vec<f32>>> {
444        let h_len = x.first().map(|r| r.len()).unwrap_or(0);
445        let default_mask = vec![false; h_len];
446        let cat_mask = cat_mask.unwrap_or(&default_mask).to_vec();
447        let out = self.predict_batch(&[x.to_vec()], &[y.to_vec()], train_size, &[cat_mask], d)?;
448        out.into_iter().next().context("predict_batch returned no batch items")
449    }
450
451    /// Runs `B` independent tables through **one** forward pass, sharing the fixed cost of the
452    /// 24-block ICL stage (and every other stage) across all of them instead of paying it once
453    /// per table. Every table must share the same row count `T` and feature count `H`, and
454    /// `train_size`/`d` are shared scalars across the whole batch — this holds for TabFM's
455    /// ensemble members, which only differ in cell *values* (per-member feature
456    /// permutation/scaling) and `cat_mask` (which position is categorical shifts with the
457    /// permutation), never in table shape. Returns `[B][T][out_dim]`.
458    pub fn predict_batch(
459        &self,
460        x_batch: &[Vec<Vec<f32>>],
461        y_batch: &[Vec<f32>],
462        train_size: usize,
463        cat_mask_batch: &[Vec<bool>],
464        d: Option<usize>,
465    ) -> Result<Vec<Vec<Vec<f32>>>> {
466        let cfg = &self.config;
467        let b_len = x_batch.len();
468        anyhow::ensure!(b_len > 0, "x_batch must have at least one table");
469        let t_len = x_batch[0].len();
470        anyhow::ensure!(t_len > 0, "x must have at least one row");
471        let h_len = x_batch[0][0].len();
472        for xb in x_batch {
473            anyhow::ensure!(xb.len() == t_len, "all batch items must have the same row count");
474            anyhow::ensure!(xb.iter().all(|r| r.len() == h_len), "all rows of x must have the same length");
475        }
476        anyhow::ensure!(y_batch.len() == b_len, "y_batch must have length B");
477        anyhow::ensure!(y_batch.iter().all(|yb| yb.len() == t_len), "each y must have the same length as x's rows");
478        anyhow::ensure!(train_size <= t_len, "train_size must be <= number of rows");
479        anyhow::ensure!(cat_mask_batch.len() == b_len, "cat_mask_batch must have length B");
480        anyhow::ensure!(cat_mask_batch.iter().all(|m| m.len() == h_len), "cat_mask must have length H");
481        let d_val = d.unwrap_or(h_len).min(h_len);
482
483        // 1. Cell embedding: [B, T, H, E]
484        let emb0 = self.cell_embed(x_batch, y_batch, train_size, cat_mask_batch, d_val)?;
485
486        // 2. Column embedding stage 1: [B, T, H, E]
487        let emb1 = self.col_embedding_forward(&emb0, train_size, &self.colenc1)?;
488
489        // 3. Prepend CLS tokens on the column axis: [B, T, row_num_cls + H, E]
490        let num_cls = cfg.row_num_cls;
491        let cls = self
492            .cls_tokens
493            .reshape((1, 1, num_cls, cfg.embed_dim))?
494            .broadcast_as((b_len, t_len, num_cls, cfg.embed_dim))?
495            .contiguous()?;
496        let emb2 = Tensor::cat(&[&cls, &emb1], 2)?;
497
498        // 4. Row interaction stage 1 (full output): [B, T, num_cls+H, E]
499        let d_plus_cls = d_val + num_cls;
500        let emb3 = self.row_interaction_forward(&emb2, d_plus_cls, &self.rowenc1, true)?;
501
502        // 5. Column embedding stage 2: [B, T, num_cls+H, E]
503        let emb4 = self.col_embedding_forward(&emb3, train_size, &self.colenc2)?;
504
505        // 6. Row interaction stage 2 (CLS-only output): [B, T, icl_dim]
506        let reps = self.row_interaction_forward(&emb4, d_plus_cls, &self.rowenc2, false)?;
507
508        // 7. In-context learning: [B, T, out_dim]
509        let logits = self.icl_forward(&reps, y_batch, train_size)?;
510
511        let out_dim = cfg.out_dim();
512        let flat: Vec<f32> = logits.flatten_all()?.to_vec1()?;
513        let mut result = vec![vec![vec![0f32; out_dim]; t_len]; b_len];
514        for (bb, batch_item) in result.iter_mut().enumerate() {
515            for (t, row) in batch_item.iter_mut().enumerate() {
516                let base = ((bb * t_len) + t) * out_dim;
517                row.copy_from_slice(&flat[base..base + out_dim]);
518            }
519        }
520        Ok(result)
521    }
522
523    // -----------------------------------------------------------------------
524    // Stage 1: cell embedding (plain Rust — grouped Fourier features per cell)
525    // -----------------------------------------------------------------------
526
527    fn cell_embed(
528        &self,
529        x_batch: &[Vec<Vec<f32>>],
530        y_batch: &[Vec<f32>],
531        train_size: usize,
532        cat_mask_batch: &[Vec<bool>],
533        d: usize,
534    ) -> Result<Tensor> {
535        let cfg = &self.config;
536        let b_len = x_batch.len();
537        let t_len = x_batch[0].len();
538        let h_len = x_batch[0][0].len();
539        let fgs = cfg.feature_group_size;
540        let num_freq = cfg.num_freq;
541        let e = cfg.embed_dim;
542        let d_safe = d.max(1);
543
544        // group index table: idx[g][h] = (h + 2^g - 1) % d_safe (shared: d is a batch-wide scalar)
545        let mut idx = vec![vec![0usize; h_len]; fgs];
546        for (g, row) in idx.iter_mut().enumerate() {
547            let offset = (1usize << g) - 1;
548            for (h, slot) in row.iter_mut().enumerate() {
549                *slot = (h + offset) % d_safe;
550            }
551        }
552
553        let ff: Vec<f32> = self.cell.fourier_freq.flatten_all()?.to_vec1()?;
554        let ffc: Vec<f32> = self.cell.fourier_freq_cat.flatten_all()?.to_vec1()?;
555        let in_w: Vec<f32> = self.cell.in_linear_w.flatten_all()?.to_vec1()?;
556        let in_b: Vec<f32> = self.cell.in_linear_b.to_vec1()?;
557        let in_w_cat: Vec<f32> = self.cell.in_linear_cat_w.flatten_all()?.to_vec1()?;
558        let in_b_cat: Vec<f32> = self.cell.in_linear_cat_b.to_vec1()?;
559
560        let y_emb_t = compute_y_embed(y_batch, &self.cell.y_embed, cfg.max_classes, &self.device)?; // [B,T,E]
561        let y_emb: Vec<f32> = y_emb_t.flatten_all()?.to_vec1()?; // [B*T*E]
562
563        let mut out = vec![0f32; b_len * t_len * h_len * e];
564        for bb in 0..b_len {
565            let x = &x_batch[bb];
566            let cat_mask = &cat_mask_batch[bb];
567            for tt in 0..t_len {
568                let add_y = tt < train_size;
569                for hh in 0..h_len {
570                    if hh >= d {
571                        continue; // padded column: leave zeroed
572                    }
573                    let mut acc = vec![0f32; e];
574                    for g in 0..fgs {
575                        let src_col = idx[g][hh];
576                        let val = x[tt][src_col];
577                        let is_cat = cat_mask.get(src_col).copied().unwrap_or(false);
578                        let (freq_row, w_lin, b_lin): (&[f32], &[f32], &[f32]) = if is_cat {
579                            (&ffc[g * num_freq..(g + 1) * num_freq], &in_w_cat, &in_b_cat)
580                        } else {
581                            (&ff[g * num_freq..(g + 1) * num_freq], &in_w, &in_b)
582                        };
583                        for ee in 0..e {
584                            let mut s = b_lin[ee];
585                            let row_off = ee * 2 * num_freq;
586                            for f in 0..num_freq {
587                                let arg = val * freq_row[f];
588                                s += arg.sin() * w_lin[row_off + f];
589                                s += arg.cos() * w_lin[row_off + num_freq + f];
590                            }
591                            acc[ee] += s;
592                        }
593                    }
594                    let base = ((bb * t_len + tt) * h_len + hh) * e;
595                    let y_base = (bb * t_len + tt) * e;
596                    for ee in 0..e {
597                        out[base + ee] = acc[ee] + if add_y { y_emb[y_base + ee] } else { 0.0 };
598                    }
599                }
600            }
601        }
602
603        Ok(Tensor::from_vec(out, (b_len, t_len, h_len, e), &self.device)?)
604    }
605
606    // -----------------------------------------------------------------------
607    // Stage 2/5: column embedding (SetTransformer, sequence axis = rows)
608    // -----------------------------------------------------------------------
609
610    fn col_embedding_forward(&self, x: &Tensor, train_size: usize, w: &ColStackWeights) -> Result<Tensor> {
611        let cfg = &self.config;
612        let (b_len, t_len, hc, e) = x.dims4()?;
613        // [B, T, HC, E] -> [B, HC, T, E] -> [(B*HC), T, E] (columns become the extended-batch
614        // axis; the same reshape-around-unchanged-attention-code trick as before, now with a
615        // real B folded into that batch axis alongside HC).
616        let src = x.permute((0, 2, 1, 3))?.contiguous()?.reshape((b_len * hc, t_len, e))?;
617
618        let mask = additive_key_mask(t_len, train_size, &self.device)?;
619
620        let mut cur = src;
621        for blk in &w.blocks {
622            let n = cur.dim(0)?;
623            let ind = blk
624                .ind_vectors
625                .unsqueeze(0)?
626                .broadcast_as((n, blk.ind_vectors.dim(0)?, blk.ind_vectors.dim(1)?))?
627                .contiguous()?;
628            let hidden = mab_forward(&ind, &cur, &cur, &blk.mab1, cfg.col_nhead, cfg.norm_eps, Some(&mask), None)?;
629            cur = mab_forward(&cur, &hidden, &hidden, &blk.mab2, cfg.col_nhead, cfg.norm_eps, None, None)?;
630        }
631
632        let projected = linear(&cur, &w.out_w_w, Some(&w.out_w_b))?;
633        let normed = rms_norm(&projected, &w.out_norm_w, cfg.norm_eps)?;
634
635        // [(B*HC), T, E] -> [B, HC, T, E] -> [B, T, HC, E]
636        let out = normed.reshape((b_len, hc, t_len, e))?.permute((0, 2, 1, 3))?.contiguous()?;
637        Ok(out)
638    }
639
640    // -----------------------------------------------------------------------
641    // Stage 4/6: row interaction (RoPE self-attention, sequence axis = columns)
642    // -----------------------------------------------------------------------
643
644    fn row_interaction_forward(
645        &self,
646        x: &Tensor, // [B, T, HC, E]
647        d_plus_cls: usize,
648        w: &RowStackWeights,
649        output_full: bool,
650    ) -> Result<Tensor> {
651        let cfg = &self.config;
652        let (b_len, t_len, hc, e) = x.dims4()?;
653        let mask = additive_key_mask(hc, d_plus_cls, &self.device)?;
654        let rope = TabfmRope { freqs: w.rope_freqs.clone() };
655
656        // [B, T, HC, E] -> [(B*T), HC, E]: B and T are both "extended batch" here (attention runs
657        // over the HC/column axis), and are already contiguous/adjacent leading dims.
658        let mut cur = x.reshape((b_len * t_len, hc, e))?;
659        for blk in &w.blocks {
660            cur = mab_forward(&cur, &cur, &cur, blk, cfg.row_nhead, cfg.norm_eps, Some(&mask), Some(&rope))?;
661        }
662
663        if output_full {
664            let normed = rms_norm(&cur, &w.out_norm_w, cfg.norm_eps)?;
665            Ok(normed.reshape((b_len, t_len, hc, e))?)
666        } else {
667            let num_cls = cfg.row_num_cls;
668            let sliced = cur.narrow(1, 0, num_cls)?; // [(B*T), num_cls, E]
669            let normed = rms_norm(&sliced, &w.out_norm_w, cfg.norm_eps)?;
670            let icl_dim = num_cls * e;
671            Ok(normed.reshape((b_len, t_len, icl_dim))?)
672        }
673    }
674
675    // -----------------------------------------------------------------------
676    // Stage 7: in-context learning
677    // -----------------------------------------------------------------------
678
679    fn icl_forward(&self, reps: &Tensor, y_batch: &[Vec<f32>], train_size: usize) -> Result<Tensor> {
680        let cfg = &self.config;
681        let (_b_len, t_len, _icl_dim) = reps.dims3()?; // reps: [B, T, icl_dim] — a real batch now,
682        // unlike before where a synthetic B=1 was faked via unsqueeze/squeeze around this stage.
683
684        let y_enc = compute_y_encoder(y_batch, &self.icl.y_encoder, cfg.max_classes, &self.device)?; // [B, T, icl_dim]
685        let tm: Vec<f32> = (0..t_len).map(|t| if t < train_size { 1.0 } else { 0.0 }).collect();
686        let tm_t = Tensor::from_vec(tm, (t_len, 1), &self.device)?;
687        let r = (reps + y_enc.broadcast_mul(&tm_t)?)?; // [B, T, icl_dim]
688
689        let mask = additive_key_mask(t_len, train_size, &self.device)?;
690
691        let mut cur = r;
692        for blk in &self.icl.blocks {
693            cur = mab_forward(&cur, &cur, &cur, blk, cfg.icl_nhead, cfg.norm_eps, Some(&mask), None)?;
694        }
695        let normed = rms_norm(&cur, &self.icl.out_norm_w, cfg.norm_eps)?;
696        mlp_forward(&normed, &self.icl.decoder)
697    }
698}
699
700// ---------------------------------------------------------------------------
701// y-embedding / y-encoder helpers
702// ---------------------------------------------------------------------------
703
704fn compute_y_embed(y_batch: &[Vec<f32>], w: &YEmbedWeights, max_classes: usize, device: &Device) -> Result<Tensor> {
705    let b_len = y_batch.len();
706    let t = y_batch[0].len();
707    match w {
708        YEmbedWeights::Embedding(emb_w) => {
709            let e = emb_w.dim(1)?;
710            let emb: Vec<f32> = emb_w.flatten_all()?.to_vec1()?;
711            let mut out = vec![0f32; b_len * t * e];
712            for (bb, y) in y_batch.iter().enumerate() {
713                for i in 0..t {
714                    let yi = y[i] as i64;
715                    if yi >= 0 && (yi as usize) < max_classes {
716                        let cls = yi as usize;
717                        let dst = (bb * t + i) * e;
718                        out[dst..dst + e].copy_from_slice(&emb[cls * e..(cls + 1) * e]);
719                    }
720                }
721            }
722            Ok(Tensor::from_vec(out, (b_len, t, e), device)?)
723        }
724        YEmbedWeights::Mlp(mlp) => {
725            let flat: Vec<f32> = y_batch.iter().flatten().copied().collect();
726            let y_col = Tensor::from_vec(flat, (b_len * t, 1), device)?;
727            let out = mlp_forward(&y_col, mlp)?;
728            let e = out.dim(1)?;
729            Ok(out.reshape((b_len, t, e))?)
730        }
731    }
732}
733
734fn compute_y_encoder(y_batch: &[Vec<f32>], w: &YEncoderWeights, max_classes: usize, device: &Device) -> Result<Tensor> {
735    let b_len = y_batch.len();
736    let t = y_batch[0].len();
737    match w {
738        YEncoderWeights::OneHot { proj_w, proj_b } => {
739            let icl_dim = proj_w.dim(0)?;
740            let w_flat: Vec<f32> = proj_w.flatten_all()?.to_vec1()?; // [icl_dim, max_classes] row-major
741            let bias: Vec<f32> = proj_b.to_vec1()?;
742            let mut out = vec![0f32; b_len * t * icl_dim];
743            for (bb, y) in y_batch.iter().enumerate() {
744                for i in 0..t {
745                    let yi = y[i] as i64;
746                    let valid = yi >= 0 && (yi as usize) < max_classes;
747                    let dst = (bb * t + i) * icl_dim;
748                    for e in 0..icl_dim {
749                        let w_contrib = if valid { w_flat[e * max_classes + yi as usize] } else { 0.0 };
750                        out[dst + e] = bias[e] + w_contrib;
751                    }
752                }
753            }
754            Ok(Tensor::from_vec(out, (b_len, t, icl_dim), device)?)
755        }
756        YEncoderWeights::Mlp(mlp) => {
757            let flat: Vec<f32> = y_batch.iter().flatten().copied().collect();
758            let y_col = Tensor::from_vec(flat, (b_len * t, 1), device)?;
759            let out = mlp_forward(&y_col, mlp)?;
760            let icl_dim = out.dim(1)?;
761            Ok(out.reshape((b_len, t, icl_dim))?)
762        }
763    }
764}
765
766// ---------------------------------------------------------------------------
767// Generic tensor ops
768// ---------------------------------------------------------------------------
769
770fn rms_norm(x: &Tensor, w: &Tensor, eps: f64) -> Result<Tensor> {
771    zsfm_nn::rms_norm(&x.to_dtype(DType::F32)?, Some(w), eps)
772}
773
774fn sigmoid(x: &Tensor) -> Result<Tensor> {
775    Ok(((x.neg()?.exp()? + 1.0)?).recip()?)
776}
777
778fn silu(x: &Tensor) -> Result<Tensor> {
779    Ok((x * sigmoid(x)?)?)
780}
781
782fn softplus(x: &Tensor) -> Result<Tensor> {
783    Ok((x.exp()? + 1.0)?.log()?)
784}
785
786fn mlp_forward(x: &Tensor, w: &MlpWeights) -> Result<Tensor> {
787    let n = w.layers.len();
788    let mut h = x.clone();
789    for (i, (lw, lb)) in w.layers.iter().enumerate() {
790        h = linear(&h, lw, Some(lb))?;
791        if i < n - 1 {
792            h = h.gelu()?;
793        }
794    }
795    Ok(h)
796}
797
798/// Additive attention mask over `[0, seq_len)` keys, valid where `key_idx < valid_len`.
799/// Shape `[1, 1, 1, seq_len]`, broadcastable against `[N, nhead, Sq, seq_len]` attention scores.
800fn additive_key_mask(seq_len: usize, valid_len: usize, device: &Device) -> Result<Tensor> {
801    let data: Vec<f32> = (0..seq_len)
802        .map(|i| if i < valid_len { 0.0 } else { -1e9 })
803        .collect();
804    Ok(Tensor::from_vec(data, (1, 1, 1, seq_len), device)?)
805}
806
807// ---------------------------------------------------------------------------
808// Interleaved-pair RoPE (checkpoint-loaded frequencies — never recomputed)
809// ---------------------------------------------------------------------------
810
811struct TabfmRope {
812    freqs: Tensor, // [head_dim/2]
813}
814
815impl TabfmRope {
816    /// Rotate `x: [N, T, nhead, head_dim]` over its `T` (dim 1) axis.
817    fn rotate(&self, x: &Tensor) -> Result<Tensor> {
818        let (_n, t, _nh, hd) = x.dims4()?;
819        let half = hd / 2;
820        let device = x.device();
821
822        let positions: Vec<f32> = (0..t).map(|p| p as f32).collect();
823        let pos = Tensor::from_vec(positions, (t,), device)?;
824        let freqs = self.freqs.to_dtype(DType::F32)?; // [half]
825        let f = pos.unsqueeze(1)?.broadcast_mul(&freqs.unsqueeze(0)?)?; // [t, half]
826        let cos = f.cos()?;
827        let sin = f.sin()?;
828        // repeat_interleave(2, -1): stack + reshape duplicates each element into adjacent pairs.
829        let cos_i = Tensor::stack(&[&cos, &cos], 2)?.reshape((t, hd))?.reshape((1, t, 1, hd))?;
830        let sin_i = Tensor::stack(&[&sin, &sin], 2)?.reshape((t, hd))?.reshape((1, t, 1, hd))?;
831
832        let x_pairs = x.reshape((x.dim(0)?, t, x.dim(2)?, half, 2))?;
833        let x1 = x_pairs.narrow(4, 0, 1)?.squeeze(4)?; // even indices
834        let x2 = x_pairs.narrow(4, 1, 1)?.squeeze(4)?; // odd indices
835        let rot = Tensor::stack(&[&x2.neg()?, &x1], 4)?.reshape(x.shape())?;
836
837        Ok((x.broadcast_mul(&cos_i)? + rot.broadcast_mul(&sin_i)?)?)
838    }
839}
840
841// ---------------------------------------------------------------------------
842// MultiheadAttentionBlock forward (attention sublayer + SwiGLU FFN sublayer)
843// ---------------------------------------------------------------------------
844
845fn mab_forward(
846    q_raw: &Tensor,
847    k_raw: &Tensor,
848    v_raw: &Tensor,
849    w: &MabWeights,
850    nhead: usize,
851    eps: f64,
852    mask: Option<&Tensor>,
853    rope: Option<&TabfmRope>,
854) -> Result<Tensor> {
855    let qn = rms_norm(q_raw, &w.pre_attn_norm, eps)?;
856    let kn = rms_norm(k_raw, &w.pre_attn_norm, eps)?;
857    let vn = rms_norm(v_raw, &w.pre_attn_norm, eps)?;
858
859    let attn_out = mha_core(&qn, &kn, &vn, w, nhead, eps, mask, rope)?;
860    let a = rms_norm(&attn_out, &w.post_attn_norm, eps)?;
861    let x = (q_raw + a)?;
862
863    let xn = rms_norm(&x, &w.pre_ff_norm, eps)?;
864    let gate = silu(&linear(&xn, &w.ffn_gate_w, Some(&w.ffn_gate_b))?)?;
865    let up = linear(&xn, &w.ffn_up_w, Some(&w.ffn_up_b))?;
866    let ff = linear(&(gate * up)?, &w.ffn_down_w, Some(&w.ffn_down_b))?;
867    let ff = rms_norm(&ff, &w.post_ff_norm, eps)?;
868
869    Ok((x + ff)?)
870}
871
872fn mha_core(
873    qn: &Tensor,
874    kn: &Tensor,
875    vn: &Tensor,
876    w: &MabWeights,
877    nhead: usize,
878    eps: f64,
879    mask: Option<&Tensor>,
880    rope: Option<&TabfmRope>,
881) -> Result<Tensor> {
882    let (n, sq, e) = qn.dims3()?;
883    let sk = kn.dim(1)?;
884    let hd = e / nhead;
885
886    let q = linear(qn, &w.q_w, Some(&w.q_b))?.reshape((n, sq, nhead, hd))?;
887    let k = linear(kn, &w.k_w, Some(&w.k_b))?.reshape((n, sk, nhead, hd))?;
888    let v = linear(vn, &w.v_w, Some(&w.v_b))?.reshape((n, sk, nhead, hd))?;
889
890    let (q, k) = match rope {
891        Some(r) => (r.rotate(&q)?, r.rotate(&k)?),
892        None => (q, k),
893    };
894
895    let q = rms_norm(&q, &w.q_norm, eps)?;
896    let k = rms_norm(&k, &w.k_norm, eps)?;
897
898    // scale = log2(e) / sqrt(hd) * softplus(per_dim_scale); attention itself uses scale=1.0.
899    let scale = (softplus(&w.per_dim_scale)? * (1.442695041_f64 / (hd as f64).sqrt()))?;
900    let q = q.broadcast_mul(&scale)?;
901
902    let q = q.permute((0, 2, 1, 3))?.contiguous()?; // [N, nhead, Sq, hd]
903    let k = k.permute((0, 2, 1, 3))?.contiguous()?;
904    let v = v.permute((0, 2, 1, 3))?.contiguous()?;
905
906    let mut scores = q.matmul(&k.transpose(D::Minus1, D::Minus2)?)?; // [N, nhead, Sq, Sk]
907    if let Some(m) = mask {
908        scores = scores.broadcast_add(m)?;
909    }
910    let probs = candle_nn::ops::softmax_last_dim(&scores)?;
911    let out = probs.matmul(&v)?; // [N, nhead, Sq, hd]
912    let out = out.permute((0, 2, 1, 3))?.contiguous()?.reshape((n, sq, e))?;
913
914    linear(&out, &w.o_w, Some(&w.o_b))
915}
916
917#[cfg(test)]
918mod tests {
919    use super::TabFMModel;
920
921    /// Compile-time check that `TabFMModel` (and every `Tensor` it holds, CPU backend) is safe
922    /// to share as `&TabFMModel` across threads — required for parallelizing the ensemble-member
923    /// loop in `ensemble::orchestrate` over a shared, read-only model reference.
924    #[test]
925    fn test_model_is_send_sync() {
926        fn assert_send_sync<T: Send + Sync>() {}
927        assert_send_sync::<TabFMModel>();
928    }
929}