Skip to main content

candle_transformers/models/z_image/
preprocess.rs

1//! Input preprocessing utilities for Z-Image
2//!
3//! Provides padding and mask construction to convert variable-length inputs
4//! into fixed-shape batch tensors.
5
6use candle::{DType, Device, Result, Tensor};
7
8use super::transformer::SEQ_MULTI_OF;
9
10/// Preprocessed inputs structure
11#[derive(Debug, Clone)]
12pub struct PreparedInputs {
13    /// Latent tensor (B, C, 1, H, W)
14    pub latents: Tensor,
15    /// Padded caption features (B, max_text_len, dim)
16    pub cap_feats: Tensor,
17    /// Caption attention mask (B, max_text_len), 1=valid, 0=padding
18    pub cap_mask: Tensor,
19    /// Original text lengths for each sample
20    pub text_lengths: Vec<usize>,
21}
22
23/// Compute padding length to align to SEQ_MULTI_OF
24#[inline]
25pub fn compute_padding_len(ori_len: usize) -> usize {
26    (SEQ_MULTI_OF - (ori_len % SEQ_MULTI_OF)) % SEQ_MULTI_OF
27}
28
29/// Pad variable-length text embeddings to uniform length
30///
31/// # Arguments
32/// * `text_embeddings` - Variable-length text embeddings, each of shape (seq_len, dim)
33/// * `pad_value` - Padding value (typically 0.0)
34/// * `device` - Device
35///
36/// # Returns
37/// * Padded tensor (B, max_len, dim)
38/// * Attention mask (B, max_len), 1=valid, 0=padding
39/// * Original lengths
40pub fn pad_text_embeddings(
41    text_embeddings: &[Tensor],
42    pad_value: f32,
43    device: &Device,
44) -> Result<(Tensor, Tensor, Vec<usize>)> {
45    if text_embeddings.is_empty() {
46        candle::bail!("text_embeddings cannot be empty");
47    }
48
49    let batch_size = text_embeddings.len();
50    let dim = text_embeddings[0].dim(1)?;
51    let dtype = text_embeddings[0].dtype();
52
53    // Compute max length and align to SEQ_MULTI_OF
54    let lengths: Vec<usize> = text_embeddings
55        .iter()
56        .map(|t| t.dim(0))
57        .collect::<Result<Vec<_>>>()?;
58    let max_len = *lengths.iter().max().unwrap();
59    let padded_len = max_len + compute_padding_len(max_len);
60
61    // Build padded tensor and mask
62    let mut padded_list = Vec::with_capacity(batch_size);
63    let mut mask_list = Vec::with_capacity(batch_size);
64
65    for (i, emb) in text_embeddings.iter().enumerate() {
66        let seq_len = lengths[i];
67        let pad_len = padded_len - seq_len;
68
69        // Pad embedding
70        let padded = if pad_len > 0 {
71            let padding = Tensor::full(pad_value, (pad_len, dim), device)?.to_dtype(dtype)?;
72            Tensor::cat(&[emb, &padding], 0)?
73        } else {
74            emb.clone()
75        };
76        padded_list.push(padded);
77
78        // Create mask: 1 for valid, 0 for padding
79        let valid = Tensor::ones((seq_len,), DType::U8, device)?;
80        let mask = if pad_len > 0 {
81            let invalid = Tensor::zeros((pad_len,), DType::U8, device)?;
82            Tensor::cat(&[&valid, &invalid], 0)?
83        } else {
84            valid
85        };
86        mask_list.push(mask);
87    }
88
89    // Stack into batch
90    let cap_feats = Tensor::stack(&padded_list, 0)?;
91    let cap_mask = Tensor::stack(&mask_list, 0)?;
92
93    Ok((cap_feats, cap_mask, lengths))
94}
95
96/// Prepare all inputs, converting variable-length inputs to fixed-shape batch tensors
97///
98/// # Arguments
99/// * `latents` - Latent tensor (B, C, H, W)
100/// * `text_embeddings` - Variable-length text embeddings, each of shape (seq_len, cap_feat_dim)
101/// * `device` - Device
102///
103/// # Returns
104/// PreparedInputs containing all preprocessed tensors
105pub fn prepare_inputs(
106    latents: &Tensor,
107    text_embeddings: &[Tensor],
108    device: &Device,
109) -> Result<PreparedInputs> {
110    // Latents: (B, C, H, W) -> (B, C, 1, H, W) add frame dimension
111    let latents = latents.unsqueeze(2)?;
112
113    // Pad text embeddings
114    let (cap_feats, cap_mask, text_lengths) = pad_text_embeddings(text_embeddings, 0.0, device)?;
115
116    Ok(PreparedInputs {
117        latents,
118        cap_feats,
119        cap_mask,
120        text_lengths,
121    })
122}
123
124/// Create attention mask for a single sample
125/// Useful for testing or simplified scenarios
126pub fn create_attention_mask(
127    valid_len: usize,
128    total_len: usize,
129    device: &Device,
130) -> Result<Tensor> {
131    let valid = Tensor::ones((valid_len,), DType::U8, device)?;
132    if valid_len < total_len {
133        let invalid = Tensor::zeros((total_len - valid_len,), DType::U8, device)?;
134        Tensor::cat(&[&valid, &invalid], 0)
135    } else {
136        Ok(valid)
137    }
138}
139
140/// Create a batch of uniform text embeddings
141///
142/// # Arguments
143/// * `text_embedding` - Single text embedding (seq_len, dim)
144/// * `batch_size` - Number of copies to create
145///
146/// # Returns
147/// Batched text embeddings (batch_size, seq_len, dim)
148pub fn batch_text_embedding(text_embedding: &Tensor, batch_size: usize) -> Result<Tensor> {
149    let (seq_len, dim) = text_embedding.dims2()?;
150    text_embedding
151        .unsqueeze(0)?
152        .broadcast_as((batch_size, seq_len, dim))?
153        .contiguous()
154}
155
156/// Create a batch of uniform masks
157///
158/// # Arguments
159/// * `mask` - Single mask (seq_len,)
160/// * `batch_size` - Number of copies to create
161///
162/// # Returns
163/// Batched masks (batch_size, seq_len)
164pub fn batch_mask(mask: &Tensor, batch_size: usize) -> Result<Tensor> {
165    let seq_len = mask.dim(0)?;
166    mask.unsqueeze(0)?
167        .broadcast_as((batch_size, seq_len))?
168        .contiguous()
169}