Skip to main content

candle_transformers/models/z_image/
text_encoder.rs

1//! Z-Image Text Encoder (Qwen3 Adapter)
2//!
3//! This module provides a Qwen3-based text encoder for Z-Image.
4//! Key difference from the standard Qwen3 model:
5//! - Returns the **second-to-last layer** hidden states (hidden_states[-2])
6//! - Does NOT apply the final RMSNorm
7
8use crate::models::with_tracing::{linear_b, Linear, RmsNorm};
9use candle::{DType, Device, Module, Result, Tensor};
10use candle_nn::{Activation, VarBuilder};
11use std::sync::Arc;
12
13/// Text Encoder configuration (Qwen3-based)
14#[derive(Debug, Clone, serde::Deserialize)]
15pub struct TextEncoderConfig {
16    #[serde(default = "default_vocab_size")]
17    pub vocab_size: usize,
18    #[serde(default = "default_hidden_size")]
19    pub hidden_size: usize,
20    #[serde(default = "default_intermediate_size")]
21    pub intermediate_size: usize,
22    #[serde(default = "default_num_hidden_layers")]
23    pub num_hidden_layers: usize,
24    #[serde(default = "default_num_attention_heads")]
25    pub num_attention_heads: usize,
26    #[serde(default = "default_num_key_value_heads")]
27    pub num_key_value_heads: usize,
28    #[serde(default = "default_head_dim")]
29    pub head_dim: usize,
30    #[serde(default = "default_rms_norm_eps")]
31    pub rms_norm_eps: f64,
32    #[serde(default = "default_rope_theta")]
33    pub rope_theta: f64,
34    #[serde(default = "default_attention_bias")]
35    pub attention_bias: bool,
36    #[serde(default = "default_hidden_act")]
37    pub hidden_act: Activation,
38    #[serde(default = "default_max_position_embeddings")]
39    pub max_position_embeddings: usize,
40}
41
42fn default_vocab_size() -> usize {
43    151936
44}
45fn default_hidden_size() -> usize {
46    2560
47}
48fn default_intermediate_size() -> usize {
49    9728
50}
51fn default_num_hidden_layers() -> usize {
52    36
53}
54fn default_num_attention_heads() -> usize {
55    32
56}
57fn default_num_key_value_heads() -> usize {
58    8
59}
60fn default_head_dim() -> usize {
61    128
62}
63fn default_rms_norm_eps() -> f64 {
64    1e-6
65}
66fn default_rope_theta() -> f64 {
67    1_000_000.0
68}
69fn default_attention_bias() -> bool {
70    false
71}
72fn default_hidden_act() -> Activation {
73    Activation::Silu
74}
75fn default_max_position_embeddings() -> usize {
76    40960
77}
78
79impl Default for TextEncoderConfig {
80    fn default() -> Self {
81        Self::z_image()
82    }
83}
84
85impl TextEncoderConfig {
86    /// Create configuration for Z-Image Text Encoder
87    pub fn z_image() -> Self {
88        Self {
89            vocab_size: 151936,
90            hidden_size: 2560,
91            intermediate_size: 9728,
92            num_hidden_layers: 36,
93            num_attention_heads: 32,
94            num_key_value_heads: 8,
95            head_dim: 128,
96            rms_norm_eps: 1e-6,
97            rope_theta: 1_000_000.0,
98            attention_bias: false,
99            hidden_act: Activation::Silu,
100            max_position_embeddings: 40960,
101        }
102    }
103}
104
105// ==================== Rotary Embedding ====================
106
107#[derive(Debug, Clone)]
108struct RotaryEmbedding {
109    sin: Tensor,
110    cos: Tensor,
111}
112
113impl RotaryEmbedding {
114    fn new(dtype: DType, cfg: &TextEncoderConfig, dev: &Device) -> Result<Self> {
115        let dim = cfg.head_dim;
116        let max_seq_len = cfg.max_position_embeddings;
117        let inv_freq: Vec<_> = (0..dim)
118            .step_by(2)
119            .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
120            .collect();
121        let inv_freq_len = inv_freq.len();
122        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(DType::F32)?;
123        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
124            .to_dtype(DType::F32)?
125            .reshape((max_seq_len, 1))?;
126        let freqs = t.matmul(&inv_freq)?;
127        Ok(Self {
128            sin: freqs.sin()?.to_dtype(dtype)?,
129            cos: freqs.cos()?.to_dtype(dtype)?,
130        })
131    }
132
133    /// Apply RoPE (q, k shape: B x H x L x D)
134    fn apply(&self, q: &Tensor, k: &Tensor, offset: usize) -> Result<(Tensor, Tensor)> {
135        let (_, _, seq_len, _) = q.dims4()?;
136        let cos = self.cos.narrow(0, offset, seq_len)?;
137        let sin = self.sin.narrow(0, offset, seq_len)?;
138        let q_embed = candle_nn::rotary_emb::rope(&q.contiguous()?, &cos, &sin)?;
139        let k_embed = candle_nn::rotary_emb::rope(&k.contiguous()?, &cos, &sin)?;
140        Ok((q_embed, k_embed))
141    }
142}
143
144// ==================== MLP ====================
145
146#[derive(Debug, Clone)]
147struct Mlp {
148    gate_proj: candle_nn::Linear,
149    up_proj: candle_nn::Linear,
150    down_proj: candle_nn::Linear,
151    act_fn: Activation,
152}
153
154impl Mlp {
155    fn new(cfg: &TextEncoderConfig, vb: VarBuilder) -> Result<Self> {
156        Ok(Self {
157            gate_proj: candle_nn::linear_no_bias(
158                cfg.hidden_size,
159                cfg.intermediate_size,
160                vb.pp("gate_proj"),
161            )?,
162            up_proj: candle_nn::linear_no_bias(
163                cfg.hidden_size,
164                cfg.intermediate_size,
165                vb.pp("up_proj"),
166            )?,
167            down_proj: candle_nn::linear_no_bias(
168                cfg.intermediate_size,
169                cfg.hidden_size,
170                vb.pp("down_proj"),
171            )?,
172            act_fn: cfg.hidden_act,
173        })
174    }
175}
176
177impl Module for Mlp {
178    fn forward(&self, x: &Tensor) -> Result<Tensor> {
179        let lhs = x.apply(&self.gate_proj)?.apply(&self.act_fn)?;
180        let rhs = x.apply(&self.up_proj)?;
181        (lhs * rhs)?.apply(&self.down_proj)
182    }
183}
184
185// ==================== Attention ====================
186
187fn repeat_kv(x: Tensor, n_rep: usize) -> Result<Tensor> {
188    if n_rep == 1 {
189        Ok(x)
190    } else {
191        let (b_sz, n_kv_head, seq_len, head_dim) = x.dims4()?;
192        x.unsqueeze(2)?
193            .broadcast_as((b_sz, n_kv_head, n_rep, seq_len, head_dim))?
194            .reshape((b_sz, n_kv_head * n_rep, seq_len, head_dim))
195    }
196}
197
198#[derive(Debug, Clone)]
199struct Attention {
200    q_proj: Linear,
201    k_proj: Linear,
202    v_proj: Linear,
203    o_proj: Linear,
204    q_norm: RmsNorm,
205    k_norm: RmsNorm,
206    num_heads: usize,
207    num_kv_heads: usize,
208    num_kv_groups: usize,
209    head_dim: usize,
210    hidden_size: usize,
211    rotary_emb: Arc<RotaryEmbedding>,
212}
213
214impl Attention {
215    fn new(
216        cfg: &TextEncoderConfig,
217        rotary_emb: Arc<RotaryEmbedding>,
218        vb: VarBuilder,
219    ) -> Result<Self> {
220        let head_dim = cfg.head_dim;
221        let num_heads = cfg.num_attention_heads;
222        let num_kv_heads = cfg.num_key_value_heads;
223        let num_kv_groups = num_heads / num_kv_heads;
224
225        let q_proj = linear_b(
226            cfg.hidden_size,
227            num_heads * head_dim,
228            cfg.attention_bias,
229            vb.pp("q_proj"),
230        )?;
231        let k_proj = linear_b(
232            cfg.hidden_size,
233            num_kv_heads * head_dim,
234            cfg.attention_bias,
235            vb.pp("k_proj"),
236        )?;
237        let v_proj = linear_b(
238            cfg.hidden_size,
239            num_kv_heads * head_dim,
240            cfg.attention_bias,
241            vb.pp("v_proj"),
242        )?;
243        let o_proj = linear_b(
244            num_heads * head_dim,
245            cfg.hidden_size,
246            cfg.attention_bias,
247            vb.pp("o_proj"),
248        )?;
249
250        let q_norm = RmsNorm::new(head_dim, cfg.rms_norm_eps, vb.pp("q_norm"))?;
251        let k_norm = RmsNorm::new(head_dim, cfg.rms_norm_eps, vb.pp("k_norm"))?;
252
253        let hidden_size = head_dim * cfg.num_attention_heads;
254
255        Ok(Self {
256            q_proj,
257            k_proj,
258            v_proj,
259            o_proj,
260            q_norm,
261            k_norm,
262            num_heads,
263            num_kv_heads,
264            num_kv_groups,
265            head_dim,
266            hidden_size,
267            rotary_emb,
268        })
269    }
270
271    fn forward(&self, x: &Tensor, attn_mask: Option<&Tensor>, offset: usize) -> Result<Tensor> {
272        let (b, l, _) = x.dims3()?;
273
274        // 1. Proj
275        let q = self.q_proj.forward(x)?;
276        let k = self.k_proj.forward(x)?;
277        let v = self.v_proj.forward(x)?;
278
279        // 2. Reshape: (B, L, H, D) -> (B, H, L, D)
280        let q = q
281            .reshape((b, l, self.num_heads, self.head_dim))?
282            .transpose(1, 2)?;
283        let k = k
284            .reshape((b, l, self.num_kv_heads, self.head_dim))?
285            .transpose(1, 2)?;
286        let v = v
287            .reshape((b, l, self.num_kv_heads, self.head_dim))?
288            .transpose(1, 2)?;
289
290        // 3. Per-head RMSNorm
291        let q_flat = q.flatten(0, 2)?;
292        let k_flat = k.flatten(0, 2)?;
293        let q_flat = self.q_norm.forward(&q_flat)?;
294        let k_flat = self.k_norm.forward(&k_flat)?;
295        let q = q_flat.reshape((b, self.num_heads, l, self.head_dim))?;
296        let k = k_flat.reshape((b, self.num_kv_heads, l, self.head_dim))?;
297
298        // 4. RoPE
299        let (q, k) = self.rotary_emb.apply(&q, &k, offset)?;
300
301        // 5. GQA repeat_kv
302        let k = repeat_kv(k, self.num_kv_groups)?.contiguous()?;
303        let v = repeat_kv(v, self.num_kv_groups)?.contiguous()?;
304
305        // 6. Attention score
306        let scale = 1.0 / (self.head_dim as f64).sqrt();
307        let mut scores = (q.matmul(&k.transpose(2, 3)?)? * scale)?;
308        if let Some(m) = attn_mask {
309            scores = scores.broadcast_add(m)?;
310        }
311        let probs = candle_nn::ops::softmax_last_dim(&scores)?;
312        let ctx = probs.matmul(&v)?; // (B, H, L, D)
313
314        // 7. Output proj
315        ctx.transpose(1, 2)?
316            .reshape((b, l, self.hidden_size))?
317            .apply(&self.o_proj)
318    }
319}
320
321// ==================== Decoder Layer ====================
322
323#[derive(Debug, Clone)]
324struct DecoderLayer {
325    self_attn: Attention,
326    mlp: Mlp,
327    ln1: RmsNorm,
328    ln2: RmsNorm,
329}
330
331impl DecoderLayer {
332    fn new(cfg: &TextEncoderConfig, rotary: Arc<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
333        let self_attn = Attention::new(cfg, rotary, vb.pp("self_attn"))?;
334        let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
335        let ln1 = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
336        let ln2 = RmsNorm::new(
337            cfg.hidden_size,
338            cfg.rms_norm_eps,
339            vb.pp("post_attention_layernorm"),
340        )?;
341        Ok(Self {
342            self_attn,
343            mlp,
344            ln1,
345            ln2,
346        })
347    }
348
349    fn forward(&self, x: &Tensor, mask: Option<&Tensor>, offset: usize) -> Result<Tensor> {
350        let h = self.ln1.forward(x)?;
351        let h = self.self_attn.forward(&h, mask, offset)?;
352        let x = (x + h)?;
353        let h2 = self.ln2.forward(&x)?;
354        let h2 = h2.apply(&self.mlp)?;
355        x + h2
356    }
357}
358
359// ==================== ZImageTextEncoder ====================
360
361/// Z-Image Text Encoder (Qwen3-based)
362///
363/// Returns the second-to-last layer hidden states (hidden_states[-2])
364/// without applying the final RMSNorm.
365#[derive(Debug, Clone)]
366pub struct ZImageTextEncoder {
367    embed_tokens: candle_nn::Embedding,
368    layers: Vec<DecoderLayer>,
369    num_hidden_layers: usize,
370    device: Device,
371    dtype: DType,
372}
373
374impl ZImageTextEncoder {
375    pub fn new(cfg: &TextEncoderConfig, vb: VarBuilder) -> Result<Self> {
376        // Note: weights have "model." prefix
377        let vb_model = vb.pp("model");
378
379        let embed_tokens =
380            candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb_model.pp("embed_tokens"))?;
381
382        let rotary = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb.device())?);
383
384        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
385        let vb_layers = vb_model.pp("layers");
386        for i in 0..cfg.num_hidden_layers {
387            layers.push(DecoderLayer::new(cfg, rotary.clone(), vb_layers.pp(i))?);
388        }
389
390        // NOTE: We do NOT load the final norm (model.norm.weight)
391        // because we return the second-to-last layer output without final norm
392
393        Ok(Self {
394            embed_tokens,
395            layers,
396            num_hidden_layers: cfg.num_hidden_layers,
397            device: vb.device().clone(),
398            dtype: vb.dtype(),
399        })
400    }
401
402    /// Create causal attention mask
403    fn causal_mask(&self, b: usize, tgt: usize, offset: usize) -> Result<Tensor> {
404        let minf = f32::NEG_INFINITY;
405        let mask: Vec<_> = (0..tgt)
406            .flat_map(|i| {
407                (0..(tgt + offset)).map(move |j| if j <= i + offset { 0.0 } else { minf })
408            })
409            .collect();
410        Tensor::from_slice(&mask, (b, 1, tgt, tgt + offset), &self.device)?.to_dtype(self.dtype)
411    }
412
413    /// Encode text, returning second-to-last layer hidden states
414    ///
415    /// # Arguments
416    /// * `input_ids` - Token IDs (B, seq_len)
417    ///
418    /// # Returns
419    /// Hidden states (B, seq_len, hidden_size) from layer[-2]
420    ///
421    /// **Important**: Returns raw output from layer[-2] WITHOUT final RMSNorm
422    pub fn forward(&self, input_ids: &Tensor) -> Result<Tensor> {
423        let (b, l) = input_ids.dims2()?;
424        let mut hidden_states = self.embed_tokens.forward(input_ids)?;
425
426        let causal = if l == 1 {
427            None
428        } else {
429            Some(self.causal_mask(b, l, 0)?)
430        };
431
432        // num_hidden_layers = 36, second-to-last layer index = 34
433        let target_layer = self.num_hidden_layers - 2;
434
435        for (i, layer) in self.layers.iter().enumerate() {
436            hidden_states = layer.forward(&hidden_states, causal.as_ref(), 0)?;
437
438            // Return after second-to-last layer, do NOT apply final norm
439            if i == target_layer {
440                return Ok(hidden_states);
441            }
442        }
443
444        // Should not reach here
445        candle::bail!("Layer index out of bounds")
446    }
447
448    /// Get the output dimension (hidden_size)
449    pub fn hidden_size(&self) -> usize {
450        // This is derived from embed_tokens weight shape
451        self.embed_tokens.embeddings().dim(1).unwrap_or(2560)
452    }
453}