1use 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#[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 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#[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 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#[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
185fn 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 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 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 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 let (q, k) = self.rotary_emb.apply(&q, &k, offset)?;
300
301 let k = repeat_kv(k, self.num_kv_groups)?.contiguous()?;
303 let v = repeat_kv(v, self.num_kv_groups)?.contiguous()?;
304
305 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)?; ctx.transpose(1, 2)?
316 .reshape((b, l, self.hidden_size))?
317 .apply(&self.o_proj)
318 }
319}
320
321#[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#[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 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 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 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 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 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 if i == target_layer {
440 return Ok(hidden_states);
441 }
442 }
443
444 candle::bail!("Layer index out of bounds")
446 }
447
448 pub fn hidden_size(&self) -> usize {
450 self.embed_tokens.embeddings().dim(1).unwrap_or(2560)
452 }
453}