1use candle_core::{D, Result, Tensor};
35use candle_nn::{self as nn, LayerNorm, LayerNormConfig, Module, VarBuilder};
36
37struct TriangleMult {
40 norm: LayerNorm,
41 left_proj: nn::Linear,
42 right_proj: nn::Linear,
43 left_gate: nn::Linear,
44 right_gate: nn::Linear,
45 out_norm: LayerNorm,
46 out_proj: nn::Linear,
47 out_gate: nn::Linear,
48 outgoing: bool,
49}
50
51impl TriangleMult {
52 fn load(vb: VarBuilder, d_pair: usize, c_hidden: usize, outgoing: bool) -> Result<Self> {
53 Ok(Self {
54 norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
55 left_proj: nn::linear_no_bias(d_pair, c_hidden, vb.pp("left_proj"))?,
56 right_proj: nn::linear_no_bias(d_pair, c_hidden, vb.pp("right_proj"))?,
57 left_gate: nn::linear_no_bias(d_pair, c_hidden, vb.pp("left_gate"))?,
58 right_gate: nn::linear_no_bias(d_pair, c_hidden, vb.pp("right_gate"))?,
59 out_norm: nn::layer_norm(c_hidden, LayerNormConfig::from(1e-5), vb.pp("out_norm"))?,
60 out_proj: nn::linear_no_bias(c_hidden, d_pair, vb.pp("out_proj"))?,
61 out_gate: nn::linear_no_bias(d_pair, d_pair, vb.pp("out_gate"))?,
62 outgoing,
63 })
64 }
65
66 fn forward(&self, z: &Tensor) -> Result<Tensor> {
67 let z_n = self.norm.forward(z)?;
68
69 let left_g = nn::ops::sigmoid(&self.left_gate.forward(&z_n)?)?;
71 let left = (self.left_proj.forward(&z_n)? * left_g)?;
72 let right_g = nn::ops::sigmoid(&self.right_gate.forward(&z_n)?)?;
73 let right = (self.right_proj.forward(&z_n)? * right_g)?;
74
75 let (left_p, right_p) = if self.outgoing {
79 (
80 left.permute((0, 3, 1, 2))?.contiguous()?,
81 right.permute((0, 3, 1, 2))?.contiguous()?,
82 )
83 } else {
84 (
85 left.permute((0, 3, 2, 1))?.contiguous()?,
86 right.permute((0, 3, 2, 1))?.contiguous()?,
87 )
88 };
89
90 let p = left_p.matmul(&right_p.transpose(D::Minus2, D::Minus1)?.contiguous()?)?; let p = p.permute((0, 2, 3, 1))?; let p = self.out_norm.forward(&p)?;
94
95 let out_g = nn::ops::sigmoid(&self.out_gate.forward(&z_n)?)?; let out = (out_g * self.out_proj.forward(&p)?)?;
97 z + &out
98 }
99}
100
101struct TriangleAttention {
104 norm: LayerNorm,
105 q_proj: nn::Linear,
106 k_proj: nn::Linear,
107 v_proj: nn::Linear,
108 pair_bias: nn::Linear,
109 gate: nn::Linear,
110 out_proj: nn::Linear,
111 n_heads: usize,
112 d_head: usize,
113 row_wise: bool,
114}
115
116impl TriangleAttention {
117 fn load(vb: VarBuilder, d_pair: usize, n_heads: usize, row_wise: bool) -> Result<Self> {
118 let d_head = d_pair / n_heads;
119 Ok(Self {
120 norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
121 q_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("q_proj"))?,
122 k_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("k_proj"))?,
123 v_proj: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("v_proj"))?,
124 pair_bias: nn::linear_no_bias(d_pair, n_heads, vb.pp("pair_bias"))?,
125 gate: nn::linear_no_bias(d_pair, n_heads * d_head, vb.pp("gate"))?,
126 out_proj: nn::linear_no_bias(n_heads * d_head, d_pair, vb.pp("out_proj"))?,
127 n_heads,
128 d_head,
129 row_wise,
130 })
131 }
132
133 fn forward_inner(&self, z: &Tensor) -> Result<Tensor> {
135 let (b, n1, n2, _) = z.dims4()?;
136 let (h, dh) = (self.n_heads, self.d_head);
137
138 let z_n = self.norm.forward(z)?;
139
140 let q = self.q_proj.forward(&z_n)?; let k = self.k_proj.forward(&z_n)?;
142 let v = self.v_proj.forward(&z_n)?;
143 let bias = self.pair_bias.forward(&z_n)?; let gate = nn::ops::sigmoid(&self.gate.forward(&z_n)?)?; let bn1 = b * n1;
148 let to_heads = |t: Tensor| -> Result<Tensor> {
149 t.reshape((bn1, n2, h, dh))?
150 .permute((0, 2, 1, 3))?
151 .contiguous()
152 };
153 let q = to_heads(q)?;
154 let k = to_heads(k)?;
155 let v = to_heads(v)?;
156
157 let bias = bias
159 .reshape((bn1, n2, h))?
160 .permute((0, 2, 1))? .unsqueeze(2)?; let scale = (dh as f64).sqrt();
164 let scores = (q.matmul(&k.transpose(D::Minus2, D::Minus1)?.contiguous()?)? / scale)?;
165 let scores = scores.broadcast_add(&bias)?;
167 let attn = nn::ops::softmax(&scores, D::Minus1)?;
168 let out = attn.matmul(&v)?; let out = out
172 .permute((0, 2, 1, 3))?
173 .contiguous()? .reshape((b, n1, n2, h * dh))?;
175
176 let out = (gate * out)?;
177 self.out_proj.forward(&out)
178 }
179
180 fn forward(&self, z: &Tensor) -> Result<Tensor> {
181 let delta = if self.row_wise {
182 self.forward_inner(z)?
183 } else {
184 let z_t = z.permute((0, 2, 1, 3))?;
186 self.forward_inner(&z_t)?.permute((0, 2, 1, 3))?
187 };
188 z + &delta
189 }
190}
191
192struct PairTransition {
195 norm: LayerNorm,
196 fc1: nn::Linear,
197 fc2: nn::Linear,
198}
199
200impl PairTransition {
201 fn load(vb: VarBuilder, d_pair: usize) -> Result<Self> {
202 Ok(Self {
203 norm: nn::layer_norm(d_pair, LayerNormConfig::from(1e-5), vb.pp("norm"))?,
204 fc1: nn::linear_no_bias(d_pair, 4 * d_pair, vb.pp("fc1"))?,
205 fc2: nn::linear_no_bias(4 * d_pair, d_pair, vb.pp("fc2"))?,
206 })
207 }
208
209 fn forward(&self, z: &Tensor) -> Result<Tensor> {
210 let h = self.norm.forward(z)?;
211 let h = self.fc1.forward(&h)?.relu()?;
212 let h = self.fc2.forward(&h)?;
213 z + &h
214 }
215}
216
217pub struct PairformerBlock {
224 tri_attn_row: TriangleAttention,
225 tri_attn_col: TriangleAttention,
226 tri_mult_out: TriangleMult,
227 tri_mult_in: TriangleMult,
228 pair_trans: PairTransition,
229}
230
231const C_HIDDEN_MULT: usize = 128;
233
234impl PairformerBlock {
235 pub fn load(vb: VarBuilder, d_pair: usize, n_heads: usize) -> Result<Self> {
239 Ok(Self {
240 tri_attn_row: TriangleAttention::load(vb.pp("tri_attn_row"), d_pair, n_heads, true)?,
241 tri_attn_col: TriangleAttention::load(vb.pp("tri_attn_col"), d_pair, n_heads, false)?,
242 tri_mult_out: TriangleMult::load(vb.pp("tri_mult_out"), d_pair, C_HIDDEN_MULT, true)?,
243 tri_mult_in: TriangleMult::load(vb.pp("tri_mult_in"), d_pair, C_HIDDEN_MULT, false)?,
244 pair_trans: PairTransition::load(vb.pp("pair_trans"), d_pair)?,
245 })
246 }
247
248 pub fn forward(&self, z: &Tensor) -> Result<Tensor> {
252 let z = self.tri_attn_row.forward(z)?;
253 let z = self.tri_attn_col.forward(&z)?;
254 let z = self.tri_mult_out.forward(&z)?;
255 let z = self.tri_mult_in.forward(&z)?;
256 self.pair_trans.forward(&z)
257 }
258}
259
260#[cfg(test)]
263mod tests {
264 use super::*;
265 use candle_core::{Device,DType, Tensor};
266
267 const B: usize = 1;
268 const N: usize = 8;
269 const D_PAIR: usize = 32; const N_HEADS: usize = 4;
271 const C_HIDDEN: usize = 16;
272
273 fn zeros_pair(device: &Device) -> Tensor {
274 Tensor::zeros(&[B, N, N, D_PAIR], DType::F32, device).unwrap()
275 }
276
277 #[test]
278 fn test_triangle_mult_outgoing_shape() {
279 let device = Device::Cpu;
280 let vb = VarBuilder::zeros(DType::F32, &device);
281 let m = TriangleMult::load(vb, D_PAIR, C_HIDDEN, true).unwrap();
282 let out = m.forward(&zeros_pair(&device)).unwrap();
283 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
284 }
285
286 #[test]
287 fn test_triangle_mult_incoming_shape() {
288 let device = Device::Cpu;
289 let vb = VarBuilder::zeros(DType::F32, &device);
290 let m = TriangleMult::load(vb, D_PAIR, C_HIDDEN, false).unwrap();
291 let out = m.forward(&zeros_pair(&device)).unwrap();
292 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
293 }
294
295 #[test]
296 fn test_triangle_attn_row_shape() {
297 let device = Device::Cpu;
298 let vb = VarBuilder::zeros(DType::F32, &device);
299 let m = TriangleAttention::load(vb, D_PAIR, N_HEADS, true).unwrap();
300 let out = m.forward(&zeros_pair(&device)).unwrap();
301 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
302 }
303
304 #[test]
305 fn test_triangle_attn_col_shape() {
306 let device = Device::Cpu;
307 let vb = VarBuilder::zeros(DType::F32, &device);
308 let m = TriangleAttention::load(vb, D_PAIR, N_HEADS, false).unwrap();
309 let out = m.forward(&zeros_pair(&device)).unwrap();
310 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
311 }
312
313 #[test]
314 fn test_pair_transition_shape() {
315 let device = Device::Cpu;
316 let vb = VarBuilder::zeros(DType::F32, &device);
317 let m = PairTransition::load(vb, D_PAIR).unwrap();
318 let out = m.forward(&zeros_pair(&device)).unwrap();
319 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
320 }
321
322 #[test]
323 fn test_pairformer_block_shape() {
324 let device = Device::Cpu;
325 let vb = VarBuilder::zeros(DType::F32, &device);
326 let block = PairformerBlock::load(vb, D_PAIR, N_HEADS).unwrap();
327 let out = block.forward(&zeros_pair(&device)).unwrap();
328 assert_eq!(out.dims(), &[B, N, N, D_PAIR]);
329 }
330
331 #[test]
332 fn test_pairformer_block_batch_shape() {
333 let device = Device::Cpu;
334 let vb = VarBuilder::zeros(DType::F32, &device);
335 let block = PairformerBlock::load(vb, D_PAIR, N_HEADS).unwrap();
336 let z = Tensor::zeros(&[2, N, N, D_PAIR], DType::F32, &device).unwrap();
337 let out = block.forward(&z).unwrap();
338 assert_eq!(out.dims(), &[2, N, N, D_PAIR]);
339 }
340
341 #[test]
342 fn test_triangle_mult_math_outgoing() {
343 let device = Device::Cpu;
346 let n = 4;
347 let d = 8;
348 let vals: Vec<f32> = (0..n * n * d).map(|i| i as f32 * 0.01).collect();
350 let z = Tensor::from_vec(vals, &[1, n, n, d], &device).unwrap();
351
352 let vb_out = VarBuilder::zeros(DType::F32, &device);
353 let vb_in = VarBuilder::zeros(DType::F32, &device);
354 let m_out = TriangleMult::load(vb_out, d, 4, true).unwrap();
355 let m_in = TriangleMult::load(vb_in, d, 4, false).unwrap();
356
357 let out_out = m_out.forward(&z).unwrap();
358 let out_in = m_in.forward(&z).unwrap();
359 assert_eq!(out_out.dims(), &[1, n, n, d]);
360 assert_eq!(out_in.dims(), &[1, n, n, d]);
361 assert_eq!(out_out.dims(), out_in.dims());
363 }
364}