candle_transformers/models/z_image/
preprocess.rs1use candle::{DType, Device, Result, Tensor};
7
8use super::transformer::SEQ_MULTI_OF;
9
10#[derive(Debug, Clone)]
12pub struct PreparedInputs {
13 pub latents: Tensor,
15 pub cap_feats: Tensor,
17 pub cap_mask: Tensor,
19 pub text_lengths: Vec<usize>,
21}
22
23#[inline]
25pub fn compute_padding_len(ori_len: usize) -> usize {
26 (SEQ_MULTI_OF - (ori_len % SEQ_MULTI_OF)) % SEQ_MULTI_OF
27}
28
29pub 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 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 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 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 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 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
96pub fn prepare_inputs(
106 latents: &Tensor,
107 text_embeddings: &[Tensor],
108 device: &Device,
109) -> Result<PreparedInputs> {
110 let latents = latents.unsqueeze(2)?;
112
113 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
124pub 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
140pub 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
156pub 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}