1use 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#[derive(Clone, Debug)]
45pub struct InferConfig {
46 embed_dim: usize,
47 max_classes: usize,
48 col_num_blocks: usize,
49 col_nhead: usize,
50 #[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 decoder_hidden: Option<usize>,
66}
67
68impl 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 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
109pub 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 pub fn config(mut self, config: InferConfig) -> Self {
138 self.config = Some(config);
139 self
140 }
141
142 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
156struct 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
179struct InducedBlockWeights {
181 ind_vectors: Tensor, 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, blocks: Vec<MabWeights>,
195 out_norm_w: Tensor,
196}
197
198struct MlpWeights {
200 layers: Vec<(Tensor, Tensor)>,
201}
202
203enum YEmbedWeights {
204 Embedding(Tensor),
206 Mlp(MlpWeights),
208}
209
210enum YEncoderWeights {
211 OneHot { proj_w: Tensor, proj_b: Tensor },
213 Mlp(MlpWeights),
215}
216
217struct CellWeights {
218 fourier_freq: Tensor, fourier_freq_cat: Tensor, in_linear_w: Tensor, in_linear_b: Tensor, in_linear_cat_w: Tensor, in_linear_cat_b: Tensor, 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, icl: IclWeights,
242}
243
244fn 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], 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 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 pub fn is_classifier(&self) -> bool {
421 self.config.is_classifier
422 }
423
424 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 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 let emb0 = self.cell_embed(x_batch, y_batch, train_size, cat_mask_batch, d_val)?;
485
486 let emb1 = self.col_embedding_forward(&emb0, train_size, &self.colenc1)?;
488
489 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 let d_plus_cls = d_val + num_cls;
500 let emb3 = self.row_interaction_forward(&emb2, d_plus_cls, &self.rowenc1, true)?;
501
502 let emb4 = self.col_embedding_forward(&emb3, train_size, &self.colenc2)?;
504
505 let reps = self.row_interaction_forward(&emb4, d_plus_cls, &self.rowenc2, false)?;
507
508 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 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 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)?; let y_emb: Vec<f32> = y_emb_t.flatten_all()?.to_vec1()?; 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; }
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 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 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 let out = normed.reshape((b_len, hc, t_len, e))?.permute((0, 2, 1, 3))?.contiguous()?;
637 Ok(out)
638 }
639
640 fn row_interaction_forward(
645 &self,
646 x: &Tensor, 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 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)?; 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 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()?; let y_enc = compute_y_encoder(y_batch, &self.icl.y_encoder, cfg.max_classes, &self.device)?; 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)?)?; 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
700fn 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()?; 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
766fn 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
798fn 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
807struct TabfmRope {
812 freqs: Tensor, }
814
815impl TabfmRope {
816 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)?; let f = pos.unsqueeze(1)?.broadcast_mul(&freqs.unsqueeze(0)?)?; let cos = f.cos()?;
827 let sin = f.sin()?;
828 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)?; let x2 = x_pairs.narrow(4, 1, 1)?.squeeze(4)?; 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
841fn 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 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()?; 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)?)?; 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)?; 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 #[test]
925 fn test_model_is_send_sync() {
926 fn assert_send_sync<T: Send + Sync>() {}
927 assert_send_sync::<TabFMModel>();
928 }
929}