ferritin_plms/esmc/layers/
geom_attention.rs1use crate::esm3::utils::affine3d::Affine3D;
2use crate::esmc::models::esmc::ESMCConfig;
3use candle_core::{D, Module, Result, Tensor};
4use candle_nn::{self as nn, LayerNorm, LayerNormConfig, Linear, VarBuilder};
5
6const SQRT_3: f64 = 1.7320508075688772;
7
8#[allow(dead_code)]
9pub struct GeometricReasoningOriginalImpl {
10 c_s: usize,
11 v_heads: usize,
12 num_vector_messages: usize,
13 mask_and_zero_frameless: bool,
14 s_norm: LayerNorm,
15 proj: Linear,
16 out_proj: Linear,
17 distance_scale_per_head: Tensor,
18 rotation_scale_per_head: Tensor,
19}
20
21impl GeometricReasoningOriginalImpl {
22 pub fn load(vb: VarBuilder, config: &ESMCConfig) -> Result<Self> {
47 let ESMCConfig {
48 d_model,
49 v_head_transformer,
50 mask_and_zero_frameless,
51 ..
52 } = config;
53
54 let num_vector_messages = 1usize;
55
56 let v_heads = v_head_transformer.unwrap_or(128);
58
59 let dim_proj = 4 * v_heads * 3 + v_heads * 3 * num_vector_messages;
60 let channels_out = v_heads * 3 * num_vector_messages;
61
62 let ln_conf = LayerNormConfig::from(1e-5);
63 let s_norm = nn::layer_norm(*d_model, ln_conf, vb.pp("layer_norm"))?;
64
65 let proj = nn::linear(*d_model, dim_proj, vb.pp("linear1"))?;
66 let out_proj = nn::linear(channels_out, *d_model, vb.pp("outproj"))?;
67 let distance_scale_per_head = Tensor::zeros((v_heads,), vb.dtype(), vb.device())?;
68 let rotation_scale_per_head = Tensor::zeros((v_heads,), vb.dtype(), vb.device())?;
69
70 Ok(Self {
71 c_s: *d_model,
72 v_heads,
73 num_vector_messages,
74 mask_and_zero_frameless: *mask_and_zero_frameless,
75 s_norm,
76 proj,
77 out_proj,
78 distance_scale_per_head,
79 rotation_scale_per_head,
80 })
81 }
82
83 pub fn forward(
93 &self,
94 s: &Tensor,
95 affine: &Affine3D,
96 affine_mask: &Tensor,
97 sequence_id: Option<&Tensor>,
98 chain_id: Option<&Tensor>,
99 ) -> Result<Tensor> {
100 let (b, l, _) = s.dims3()?;
101 let dtype = s.dtype();
102 let device = s.device();
103
104 let attn_bias = if let Some(seq_id) = sequence_id {
107 let seq_q = seq_id.unsqueeze(D::Minus1)?; let seq_k = seq_id.unsqueeze(D::Minus2)?; let same_seq = seq_q
111 .broadcast_as((b, l, l))?
112 .eq(&seq_k.broadcast_as((b, l, l))?)?
113 .to_dtype(dtype)?;
114 same_seq.unsqueeze(1)? } else {
116 Tensor::ones((b, 1, l, l), dtype, device)?
117 };
118
119 let neg_inf =
121 Tensor::full(f32::NEG_INFINITY as f64, attn_bias.shape(), device)?.to_dtype(dtype)?;
122 let frame_mask_k = affine_mask
124 .unsqueeze(1)?
125 .unsqueeze(1)?
126 .broadcast_as(attn_bias.shape())?;
127 let mut attn_bias = frame_mask_k.where_cond(&attn_bias, &neg_inf)?;
128
129 if let Some(cid) = chain_id {
131 let chain_q = cid.unsqueeze(D::Minus1)?.broadcast_as((b, l, l))?;
132 let chain_k = cid.unsqueeze(D::Minus2)?.broadcast_as((b, l, l))?;
133 let diff_chain = chain_q.ne(&chain_k)?.unsqueeze(1)?; let diff_bias = diff_chain.broadcast_as(attn_bias.shape())?;
135 attn_bias = diff_bias.where_cond(&neg_inf, &attn_bias)?;
136 }
137
138 let ns = self.s_norm.forward(s)?;
140 let proj_out = self.proj.forward(&ns)?;
141
142 let vec_rot_size = self.v_heads * 2 * 3 + self.v_heads * 3 * self.num_vector_messages;
143 let vec_dist_size = self.v_heads * 2 * 3;
144 let vec_rot = proj_out.narrow(D::Minus1, 0, vec_rot_size)?;
145 let vec_dist = proj_out.narrow(D::Minus1, vec_rot_size, vec_dist_size)?;
146
147 let h_rot = 2 * self.v_heads + self.v_heads * self.num_vector_messages;
150 let vec_rot = vec_rot.reshape((b, l, h_rot, 3))?;
151 let vec_rot = Affine3D::apply_rot(&affine.rot, &vec_rot)?;
153
154 let query_rot = vec_rot.narrow(D::Minus2, 0, self.v_heads)?; let key_rot = vec_rot.narrow(D::Minus2, self.v_heads, self.v_heads)?;
156 let value = vec_rot.narrow(
157 D::Minus2,
158 2 * self.v_heads,
159 self.v_heads * self.num_vector_messages,
160 )?;
161
162 let vec_dist = vec_dist.reshape((b, l, self.v_heads * 2, 3))?;
164 let vec_dist = affine.apply(&vec_dist)?; let query_dist = vec_dist.narrow(D::Minus2, 0, self.v_heads)?; let key_dist = vec_dist.narrow(D::Minus2, self.v_heads, self.v_heads)?;
167
168 let query_rot = query_rot.permute((0, 2, 1, 3))?;
171 let key_rot = key_rot.permute((0, 2, 3, 1))?;
173 let query_dist = query_dist.permute((0, 2, 1, 3))?.unsqueeze(D::Minus2)?;
175 let key_dist = key_dist.permute((0, 2, 1, 3))?.unsqueeze(2)?;
177 let value = value
179 .reshape((b, l, self.v_heads, self.num_vector_messages * 3))?
180 .permute((0, 2, 1, 3))?;
181
182 let rotation_term = query_rot
186 .contiguous()?
187 .matmul(&key_rot.contiguous()?)?
188 .affine(SQRT_3.recip(), 0.0)?;
189
190 let diff = query_dist
192 .broadcast_as((b, self.v_heads, l, l, 3))?
193 .sub(&key_dist.broadcast_as((b, self.v_heads, l, l, 3))?)?;
194 let distance_term = diff
195 .sqr()?
196 .sum(D::Minus1)?
197 .sqrt()?
198 .affine(SQRT_3.recip(), 0.0)?;
199
200 let dist_w = softplus(&self.distance_scale_per_head)?
202 .reshape((1, self.v_heads, 1, 1))?
203 .broadcast_as((b, self.v_heads, l, l))?;
204 let rot_w = softplus(&self.rotation_scale_per_head)?
205 .reshape((1, self.v_heads, 1, 1))?
206 .broadcast_as((b, self.v_heads, l, l))?;
207
208 let mut attn_weight = rotation_term
209 .mul(&rot_w)?
210 .sub(&distance_term.mul(&dist_w)?)?;
211
212 let attn_bias = attn_bias.broadcast_as(attn_weight.shape())?;
214 attn_weight = attn_weight.add(&attn_bias)?;
215
216 let attn_weight = candle_nn::ops::softmax(&attn_weight, D::Minus1)?;
217
218 let attn_out = attn_weight.matmul(&value.contiguous()?)?;
221
222 let attn_out = attn_out
224 .permute((0, 2, 1, 3))? .contiguous()?
226 .reshape((b, l, self.v_heads * self.num_vector_messages, 3))?;
227
228 let attn_out = Affine3D::apply_rot_inv(&affine.rot, &attn_out)?;
230
231 let mut attn_out =
233 attn_out
234 .contiguous()?
235 .reshape((b, l, self.v_heads * self.num_vector_messages * 3))?;
236
237 if self.mask_and_zero_frameless {
239 let zeros = Tensor::zeros_like(&attn_out)?;
240 let mask_exp = affine_mask
241 .unsqueeze(D::Minus1)?
242 .broadcast_as(attn_out.shape())?;
243 attn_out = mask_exp.where_cond(&attn_out, &zeros)?;
244 }
245
246 self.out_proj.forward(&attn_out)
247 }
248}
249
250fn softplus(t: &Tensor) -> Result<Tensor> {
252 t.exp()?.affine(1.0, 1.0)?.log()
254}