Skip to main content

candle_transformers/models/
lfm2.rs

1//! LFM2 (Liquid Foundation Model 2) implementation.
2//!
3//! LFM2 is a hybrid architecture that combines attention and short convolution layers.
4//! See [LiquidAI](https://www.liquid.ai/) for more information.
5//!
6//! This implementation supports the LFM2ForCausalLM architecture from HuggingFace transformers.
7
8use crate::models::with_tracing::{linear_no_bias as linear, Embedding, Linear, RmsNorm};
9use crate::utils::repeat_kv;
10use candle::{DType, Device, IndexOp, Module, Result, Tensor};
11use candle_nn::{Conv1d, Conv1dConfig, VarBuilder};
12use std::collections::HashMap;
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Deserialize)]
15#[serde(rename_all = "snake_case")]
16pub enum LayerType {
17    FullAttention,
18    Conv,
19}
20
21#[derive(Debug, Clone, serde::Deserialize)]
22pub struct Lfm2Config {
23    pub vocab_size: usize,
24    pub hidden_size: usize,
25    pub num_hidden_layers: usize,
26    pub num_attention_heads: usize,
27    #[serde(default = "default_num_key_value_heads")]
28    pub num_key_value_heads: usize,
29    #[serde(default = "default_norm_eps")]
30    pub norm_eps: f64,
31    #[serde(default = "default_rope_theta")]
32    pub rope_theta: f32,
33    #[serde(default = "default_max_position_embeddings")]
34    pub max_position_embeddings: usize,
35    #[serde(default = "default_conv_l_cache", alias = "conv_L_cache")]
36    pub conv_l_cache: usize,
37    #[serde(default)]
38    pub conv_bias: bool,
39    pub layer_types: Vec<LayerType>,
40    #[serde(default)]
41    pub tie_embedding: bool,
42    pub bos_token_id: Option<u32>,
43    pub eos_token_id: Option<u32>,
44    // FFN dimension configuration
45    #[serde(default = "default_ffn_dim_multiplier")]
46    pub block_ffn_dim_multiplier: f32,
47    #[serde(default = "default_block_multiple_of")]
48    pub block_multiple_of: usize,
49}
50
51fn default_num_key_value_heads() -> usize {
52    8
53}
54
55fn default_norm_eps() -> f64 {
56    1e-5
57}
58
59fn default_rope_theta() -> f32 {
60    1_000_000.0
61}
62
63fn default_max_position_embeddings() -> usize {
64    128000
65}
66
67fn default_conv_l_cache() -> usize {
68    3
69}
70
71fn default_ffn_dim_multiplier() -> f32 {
72    1.0
73}
74
75fn default_block_multiple_of() -> usize {
76    256
77}
78
79impl Lfm2Config {
80    pub fn head_dim(&self) -> usize {
81        self.hidden_size / self.num_attention_heads
82    }
83
84    /// Compute the actual intermediate size for the FFN.
85    /// LFM2 uses: hidden_size * 4 * block_ffn_dim_multiplier, rounded to block_multiple_of
86    fn compute_intermediate_size(&self) -> usize {
87        let base_size = (self.hidden_size as f32 * 4.0 * self.block_ffn_dim_multiplier) as usize;
88        let multiple = self.block_multiple_of;
89        base_size.div_ceil(multiple) * multiple
90    }
91
92    pub fn into_config(self, use_flash_attn: bool) -> Config {
93        // Use computed intermediate size (matches actual weights) instead of config field
94        let intermediate_size = self.compute_intermediate_size();
95        Config {
96            vocab_size: self.vocab_size,
97            hidden_size: self.hidden_size,
98            intermediate_size,
99            num_hidden_layers: self.num_hidden_layers,
100            num_attention_heads: self.num_attention_heads,
101            num_key_value_heads: self.num_key_value_heads,
102            norm_eps: self.norm_eps,
103            rope_theta: self.rope_theta,
104            max_position_embeddings: self.max_position_embeddings,
105            conv_l_cache: self.conv_l_cache,
106            conv_bias: self.conv_bias,
107            layer_types: self.layer_types,
108            tie_embedding: self.tie_embedding,
109            bos_token_id: self.bos_token_id,
110            eos_token_id: self.eos_token_id,
111            use_flash_attn,
112        }
113    }
114}
115
116#[derive(Debug, Clone)]
117pub struct Config {
118    pub vocab_size: usize,
119    pub hidden_size: usize,
120    pub intermediate_size: usize,
121    pub num_hidden_layers: usize,
122    pub num_attention_heads: usize,
123    pub num_key_value_heads: usize,
124    pub norm_eps: f64,
125    pub rope_theta: f32,
126    pub max_position_embeddings: usize,
127    pub conv_l_cache: usize,
128    pub conv_bias: bool,
129    pub layer_types: Vec<LayerType>,
130    pub tie_embedding: bool,
131    pub bos_token_id: Option<u32>,
132    pub eos_token_id: Option<u32>,
133    pub use_flash_attn: bool,
134}
135
136impl Config {
137    pub fn head_dim(&self) -> usize {
138        self.hidden_size / self.num_attention_heads
139    }
140}
141
142/// Cache for LFM2 model supporting both attention KV cache and convolution state cache.
143#[derive(Debug, Clone)]
144pub struct Cache {
145    masks: HashMap<(usize, usize), Tensor>,
146    pub use_kv_cache: bool,
147    // KV cache for attention layers: (key, value) per layer
148    kvs: Vec<Option<(Tensor, Tensor)>>,
149    // Conv state cache for convolution layers
150    conv_states: Vec<Option<Tensor>>,
151    cos: Tensor,
152    sin: Tensor,
153    device: Device,
154}
155
156fn calculate_default_inv_freq(cfg: &Config) -> Vec<f32> {
157    let head_dim = cfg.head_dim();
158    (0..head_dim)
159        .step_by(2)
160        .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / head_dim as f32))
161        .collect()
162}
163
164impl Cache {
165    pub fn new(use_kv_cache: bool, dtype: DType, config: &Config, device: &Device) -> Result<Self> {
166        let theta = calculate_default_inv_freq(config);
167        let theta = Tensor::new(theta, device)?;
168
169        let idx_theta = Tensor::arange(0, config.max_position_embeddings as u32, device)?
170            .to_dtype(DType::F32)?
171            .reshape((config.max_position_embeddings, 1))?
172            .matmul(&theta.reshape((1, theta.elem_count()))?)?;
173        let cos = idx_theta.cos()?.to_dtype(dtype)?;
174        let sin = idx_theta.sin()?.to_dtype(dtype)?;
175
176        let num_layers = config.num_hidden_layers;
177        Ok(Self {
178            masks: HashMap::new(),
179            use_kv_cache,
180            kvs: vec![None; num_layers],
181            conv_states: vec![None; num_layers],
182            device: device.clone(),
183            cos,
184            sin,
185        })
186    }
187
188    fn mask(&mut self, seq_len: usize, index_pos: usize) -> Result<Tensor> {
189        let kv_len = index_pos + seq_len;
190        if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
191            Ok(mask.clone())
192        } else {
193            let mask = crate::utils::build_causal_mask(seq_len, index_pos, &self.device)?;
194            self.masks.insert((seq_len, kv_len), mask.clone());
195            Ok(mask)
196        }
197    }
198
199    pub fn clear(&mut self) {
200        self.kvs.iter_mut().for_each(|v| *v = None);
201        self.conv_states.iter_mut().for_each(|v| *v = None);
202    }
203}
204
205fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
206    let shape = mask.shape();
207    let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
208    let m = mask.where_cond(&on_true, on_false)?;
209    Ok(m)
210}
211
212#[cfg(feature = "flash-attn")]
213fn flash_attn(
214    q: &Tensor,
215    k: &Tensor,
216    v: &Tensor,
217    softmax_scale: f32,
218    causal: bool,
219) -> Result<Tensor> {
220    candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
221}
222
223#[cfg(not(feature = "flash-attn"))]
224fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
225    unimplemented!("compile with '--features flash-attn'")
226}
227
228/// MLP layer with SwiGLU activation.
229#[derive(Debug, Clone)]
230struct Mlp {
231    gate_proj: Linear,
232    up_proj: Linear,
233    down_proj: Linear,
234    span: tracing::Span,
235}
236
237impl Mlp {
238    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
239        let hidden_size = cfg.hidden_size;
240        let intermediate_size = cfg.intermediate_size;
241        // LFM2 uses w1 (gate), w3 (up), w2 (down) naming convention
242        let gate_proj = linear(hidden_size, intermediate_size, vb.pp("w1"))?;
243        let up_proj = linear(hidden_size, intermediate_size, vb.pp("w3"))?;
244        let down_proj = linear(intermediate_size, hidden_size, vb.pp("w2"))?;
245        Ok(Self {
246            gate_proj,
247            up_proj,
248            down_proj,
249            span: tracing::span!(tracing::Level::TRACE, "mlp"),
250        })
251    }
252
253    fn forward(&self, x: &Tensor) -> Result<Tensor> {
254        let _enter = self.span.enter();
255        let gate = candle_nn::ops::silu(&self.gate_proj.forward(x)?)?;
256        let up = self.up_proj.forward(x)?;
257        self.down_proj.forward(&(gate * up)?)
258    }
259}
260
261/// Attention layer with per-head QK normalization and RoPE.
262#[derive(Debug, Clone)]
263struct Attention {
264    q_proj: Linear,
265    k_proj: Linear,
266    v_proj: Linear,
267    o_proj: Linear,
268    q_norm: RmsNorm,
269    k_norm: RmsNorm,
270    num_attention_heads: usize,
271    num_key_value_heads: usize,
272    head_dim: usize,
273    use_flash_attn: bool,
274    span: tracing::Span,
275    span_rot: tracing::Span,
276}
277
278impl Attention {
279    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
280        let hidden_size = cfg.hidden_size;
281        let num_attention_heads = cfg.num_attention_heads;
282        let num_key_value_heads = cfg.num_key_value_heads;
283        let head_dim = cfg.head_dim();
284
285        let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
286        let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
287        let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
288        let o_proj = linear(
289            num_attention_heads * head_dim,
290            hidden_size,
291            vb.pp("out_proj"),
292        )?;
293
294        let q_norm = RmsNorm::new(head_dim, cfg.norm_eps, vb.pp("q_layernorm"))?;
295        let k_norm = RmsNorm::new(head_dim, cfg.norm_eps, vb.pp("k_layernorm"))?;
296
297        Ok(Self {
298            q_proj,
299            k_proj,
300            v_proj,
301            o_proj,
302            q_norm,
303            k_norm,
304            num_attention_heads,
305            num_key_value_heads,
306            head_dim,
307            use_flash_attn: cfg.use_flash_attn,
308            span: tracing::span!(tracing::Level::TRACE, "attn"),
309            span_rot: tracing::span!(tracing::Level::TRACE, "attn-rot"),
310        })
311    }
312
313    fn apply_rotary_emb(&self, x: &Tensor, index_pos: usize, cache: &Cache) -> Result<Tensor> {
314        let _enter = self.span_rot.enter();
315        let (_, _, seq_len, _) = x.dims4()?;
316        let cos = cache.cos.narrow(0, index_pos, seq_len)?;
317        let sin = cache.sin.narrow(0, index_pos, seq_len)?;
318        candle_nn::rotary_emb::rope(&x.contiguous()?, &cos, &sin)
319    }
320
321    fn forward(
322        &self,
323        x: &Tensor,
324        index_pos: usize,
325        block_idx: usize,
326        cache: &mut Cache,
327    ) -> Result<Tensor> {
328        let _enter = self.span.enter();
329        let (b_sz, seq_len, _) = x.dims3()?;
330
331        let q = self.q_proj.forward(x)?;
332        let k = self.k_proj.forward(x)?;
333        let v = self.v_proj.forward(x)?;
334
335        // Reshape to (batch, seq, num_heads, head_dim) then transpose to (batch, num_heads, seq, head_dim)
336        let q = q
337            .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?
338            .transpose(1, 2)?;
339        let k = k
340            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
341            .transpose(1, 2)?;
342        let v = v
343            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
344            .transpose(1, 2)?
345            .contiguous()?;
346
347        // Apply per-head QK normalization
348        let q = self.q_norm.forward(&q.contiguous()?)?;
349        let k = self.k_norm.forward(&k.contiguous()?)?;
350
351        // Apply rotary embeddings
352        let q = self.apply_rotary_emb(&q, index_pos, cache)?;
353        let k = self.apply_rotary_emb(&k, index_pos, cache)?;
354
355        // Handle KV cache
356        let (k, v) = if cache.use_kv_cache {
357            match &cache.kvs[block_idx] {
358                Some((k_cache, v_cache)) if index_pos > 0 => {
359                    let k = Tensor::cat(&[k_cache, &k], 2)?.contiguous()?;
360                    let v = Tensor::cat(&[v_cache, &v], 2)?.contiguous()?;
361                    (k, v)
362                }
363                _ => (k, v),
364            }
365        } else {
366            (k, v)
367        };
368
369        if cache.use_kv_cache {
370            cache.kvs[block_idx] = Some((k.clone(), v.clone()));
371        }
372
373        // Expand KV heads to match query heads
374        let k = repeat_kv(k, self.num_attention_heads / self.num_key_value_heads)?;
375        let v = repeat_kv(v, self.num_attention_heads / self.num_key_value_heads)?;
376
377        let y = if self.use_flash_attn {
378            let q = q.transpose(1, 2)?;
379            let k = k.transpose(1, 2)?;
380            let v = v.transpose(1, 2)?;
381            let softmax_scale = 1f32 / (self.head_dim as f32).sqrt();
382            flash_attn(&q, &k, &v, softmax_scale, seq_len > 1)?.transpose(1, 2)?
383        } else {
384            let in_dtype = q.dtype();
385            let q = q.to_dtype(DType::F32)?;
386            let k = k.to_dtype(DType::F32)?;
387            let v = v.to_dtype(DType::F32)?;
388            let att = (q.matmul(&k.t()?)? / (self.head_dim as f64).sqrt())?;
389            let att = if seq_len == 1 {
390                att
391            } else {
392                let mask = cache.mask(seq_len, index_pos)?.broadcast_as(att.shape())?;
393                masked_fill(&att, &mask, f32::NEG_INFINITY)?
394            };
395            let att = candle_nn::ops::softmax_last_dim(&att)?;
396            att.matmul(&v.contiguous()?)?.to_dtype(in_dtype)?
397        };
398
399        let y = y.transpose(1, 2)?.reshape((
400            b_sz,
401            seq_len,
402            self.num_attention_heads * self.head_dim,
403        ))?;
404        self.o_proj.forward(&y)
405    }
406}
407
408/// Short convolution layer for efficient sequence processing.
409#[derive(Debug, Clone)]
410struct ShortConv {
411    in_proj: Linear,
412    out_proj: Linear,
413    conv_weight: Tensor,
414    l_cache: usize,
415    hidden_size: usize,
416    span: tracing::Span,
417}
418
419impl ShortConv {
420    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
421        let hidden_size = cfg.hidden_size;
422        let l_cache = cfg.conv_l_cache;
423
424        // in_proj projects to 3 * hidden_size for B, C, X components
425        let in_proj = linear(hidden_size, 3 * hidden_size, vb.pp("in_proj"))?;
426        let out_proj = linear(hidden_size, hidden_size, vb.pp("out_proj"))?;
427
428        // Conv weight shape: (hidden_size, 1, l_cache) or (hidden_size, l_cache)
429        let conv_weight = vb.get((hidden_size, 1, l_cache), "conv.weight")?;
430
431        Ok(Self {
432            in_proj,
433            out_proj,
434            conv_weight,
435            l_cache,
436            hidden_size,
437            span: tracing::span!(tracing::Level::TRACE, "shortconv"),
438        })
439    }
440
441    fn forward(&self, x: &Tensor, block_idx: usize, cache: &mut Cache) -> Result<Tensor> {
442        let _enter = self.span.enter();
443        let (b_sz, seq_len, _) = x.dims3()?;
444
445        // Project input to B, C, X components
446        let bcx = self.in_proj.forward(x)?.transpose(1, 2)?;
447        let b = bcx.narrow(1, 0, self.hidden_size)?;
448        let c = bcx.narrow(1, self.hidden_size, self.hidden_size)?;
449        let x_proj = bcx.narrow(1, 2 * self.hidden_size, self.hidden_size)?;
450
451        // Element-wise multiply B and X
452        let bx = (b * &x_proj)?.contiguous()?;
453
454        // Prepare conv weight: squeeze to (hidden_size, l_cache) for element-wise, or keep for Conv1d
455        let conv_weight = self.conv_weight.squeeze(1)?;
456
457        let conv_out = if seq_len == 1 {
458            // Token-by-token generation: use cached state
459            let mut state = match &cache.conv_states[block_idx] {
460                Some(s) => s.clone(),
461                None => Tensor::zeros(
462                    (b_sz, self.hidden_size, self.l_cache),
463                    bx.dtype(),
464                    bx.device(),
465                )?,
466            };
467
468            // Shift cache and add new token
469            if self.l_cache > 1 {
470                let tail = state.narrow(2, 1, self.l_cache - 1)?;
471                state = Tensor::cat(&[tail, bx.clone()], 2)?;
472            } else {
473                state = bx.clone();
474            }
475
476            if cache.use_kv_cache {
477                cache.conv_states[block_idx] = Some(state.clone());
478            }
479
480            // Apply convolution as element-wise multiply and sum
481            (state * conv_weight.unsqueeze(0)?)?
482                .sum_keepdim(2)?
483                .contiguous()?
484        } else {
485            // Prefill: use Conv1d
486            let conv = Conv1d::new(
487                self.conv_weight.clone(),
488                None,
489                Conv1dConfig {
490                    padding: self.l_cache.saturating_sub(1),
491                    groups: self.hidden_size,
492                    ..Default::default()
493                },
494            );
495            let mut out = conv.forward(&bx)?;
496            out = out.narrow(2, 0, seq_len)?;
497
498            // Update cache with last l_cache tokens
499            if cache.use_kv_cache && self.l_cache > 0 {
500                let start = seq_len.saturating_sub(self.l_cache);
501                let cache_len = seq_len - start;
502                let mut cache_src = bx.narrow(2, start, cache_len)?;
503                if cache_len < self.l_cache {
504                    let pad = self.l_cache - cache_len;
505                    let zeros = Tensor::zeros(
506                        (b_sz, self.hidden_size, pad),
507                        cache_src.dtype(),
508                        cache_src.device(),
509                    )?;
510                    cache_src = Tensor::cat(&[zeros, cache_src], 2)?;
511                }
512                cache.conv_states[block_idx] = Some(cache_src);
513            }
514
515            out
516        };
517
518        // Multiply by C and project output
519        let conv_out = (c * &conv_out)?;
520        let conv_out = conv_out.transpose(1, 2)?.contiguous()?;
521        self.out_proj.forward(&conv_out)
522    }
523}
524
525/// Unified decoder layer supporting both attention and convolution.
526#[derive(Debug, Clone)]
527enum LayerKind {
528    Attention(Box<Attention>),
529    ShortConv(ShortConv),
530}
531
532#[derive(Debug, Clone)]
533struct DecoderLayer {
534    input_layernorm: RmsNorm,
535    post_attention_layernorm: RmsNorm,
536    mlp: Mlp,
537    kind: LayerKind,
538    span: tracing::Span,
539}
540
541impl DecoderLayer {
542    fn new(cfg: &Config, layer_idx: usize, vb: VarBuilder) -> Result<Self> {
543        // LFM2 uses operator_norm and ffn_norm naming
544        let input_layernorm = RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb.pp("operator_norm"))?;
545        let post_attention_layernorm =
546            RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb.pp("ffn_norm"))?;
547        // LFM2 uses feed_forward naming for MLP
548        let mlp = Mlp::new(cfg, vb.pp("feed_forward"))?;
549
550        let layer_type = cfg
551            .layer_types
552            .get(layer_idx)
553            .copied()
554            .unwrap_or(LayerType::FullAttention);
555        let kind = match layer_type {
556            LayerType::FullAttention => {
557                LayerKind::Attention(Box::new(Attention::new(cfg, vb.pp("self_attn"))?))
558            }
559            LayerType::Conv => LayerKind::ShortConv(ShortConv::new(cfg, vb.pp("conv"))?),
560        };
561
562        Ok(Self {
563            input_layernorm,
564            post_attention_layernorm,
565            mlp,
566            kind,
567            span: tracing::span!(tracing::Level::TRACE, "layer"),
568        })
569    }
570
571    fn forward(
572        &self,
573        x: &Tensor,
574        index_pos: usize,
575        block_idx: usize,
576        cache: &mut Cache,
577    ) -> Result<Tensor> {
578        let _enter = self.span.enter();
579        let residual = x;
580        let x = self.input_layernorm.forward(x)?;
581
582        let x = match &self.kind {
583            LayerKind::Attention(attn) => attn.forward(&x, index_pos, block_idx, cache)?,
584            LayerKind::ShortConv(conv) => conv.forward(&x, block_idx, cache)?,
585        };
586
587        let x = (x + residual)?;
588        let residual = &x;
589        let x = self.post_attention_layernorm.forward(&x)?;
590        let x = self.mlp.forward(&x)?;
591        x + residual
592    }
593}
594
595/// LFM2 model for causal language modeling.
596#[derive(Debug, Clone)]
597pub struct Model {
598    embed_tokens: Embedding,
599    layers: Vec<DecoderLayer>,
600    embedding_norm: RmsNorm,
601    lm_head: Linear,
602    dtype: DType,
603}
604
605impl Model {
606    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
607        let vb_m = vb.pp("model");
608
609        let embed_tokens =
610            Embedding::new(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
611
612        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
613        let vb_l = vb_m.pp("layers");
614        for layer_idx in 0..cfg.num_hidden_layers {
615            let layer = DecoderLayer::new(cfg, layer_idx, vb_l.pp(layer_idx))?;
616            layers.push(layer);
617        }
618
619        let embedding_norm =
620            RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb_m.pp("embedding_norm"))?;
621
622        let lm_head = if cfg.tie_embedding {
623            Linear::from_weights(embed_tokens.embeddings().clone(), None)
624        } else {
625            linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
626        };
627
628        Ok(Self {
629            embed_tokens,
630            layers,
631            embedding_norm,
632            lm_head,
633            dtype: vb.dtype(),
634        })
635    }
636
637    pub fn forward(
638        &self,
639        input_ids: &Tensor,
640        index_pos: usize,
641        cache: &mut Cache,
642    ) -> Result<Tensor> {
643        let (_, seq_len) = input_ids.dims2()?;
644        let mut hidden_states = self.embed_tokens.forward(input_ids)?;
645
646        for (block_idx, layer) in self.layers.iter().enumerate() {
647            hidden_states = layer.forward(&hidden_states, index_pos, block_idx, cache)?;
648        }
649
650        let hidden_states = self.embedding_norm.forward(&hidden_states)?;
651        let hidden_states = hidden_states.i((.., seq_len - 1, ..))?.contiguous()?;
652        let logits = self.lm_head.forward(&hidden_states)?;
653        logits.to_dtype(DType::F32)
654    }
655
656    pub fn dtype(&self) -> DType {
657        self.dtype
658    }
659}