Skip to main content

ferritin_plms/esmc/layers/
geom_attention.rs

1use 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 new(
23    //     c_s: i64,
24    //     v_heads: i64,
25    //     num_vector_messages: i64,
26    //     mask_and_zero_frameless: bool,
27    //     _divide_residual_by_depth: bool,
28    //     bias: bool,
29    //     device: &Device,
30    // ) -> Result<Self> {
31    //     let dim_proj = 4 * v_heads * 3 + v_heads * 3 * num_vector_messages;
32    //     let channels_out = v_heads * 3 * num_vector_messages;
33
34    //     Ok(Self {
35    //         c_s,
36    //         v_heads,
37    //         num_vector_messages,
38    //         mask_and_zero_frameless,
39    //         s_norm: LayerNorm::new(c_s, bias)?,
40    //         proj: Linear::new(c_s, dim_proj, bias)?,
41    //         out_proj: Linear::new(channels_out, c_s, bias)?,
42    //         distance_scale_per_head: Tensor::zeros((v_heads,), device)?,
43    //         rotation_scale_per_head: Tensor::zeros((v_heads,), device)?,
44    //     })
45    // }
46    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        // todo: this is a hidden param. Needs to be fixed
57        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    /// Geometric attention forward pass.
84    ///
85    /// - `s`:           `(B, L, d_model)` hidden states.
86    /// - `affine`:      per-residue local frames `(B, L, 3, 3)` rot + `(B, L, 3)` trans.
87    /// - `affine_mask`: `(B, L)` u8 — 1 where the frame is valid, 0 for frameless positions.
88    /// - `sequence_id`: optional `(B, L)` int — positions with the same ID form one protein.
89    /// - `chain_id`:    optional `(B, L)` int — positions in the same chain.
90    ///
91    /// Returns `(B, L, d_model)`.
92    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        // ── Attention bias (sequence_id and chain_id masking) ──────────────
105        // Same-sequence pairs get 1.0; cross-sequence and frameless get -inf.
106        let attn_bias = if let Some(seq_id) = sequence_id {
107            let seq_q = seq_id.unsqueeze(D::Minus1)?; // (B, L, 1)
108            let seq_k = seq_id.unsqueeze(D::Minus2)?; // (B, 1, L)
109            // (B, L, L): 1 where same sequence, 0 where different
110            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)? // (B, 1, L, L)
115        } else {
116            Tensor::ones((b, 1, l, l), dtype, device)?
117        };
118
119        // Mask frameless key positions with -inf
120        let neg_inf =
121            Tensor::full(f32::NEG_INFINITY as f64, attn_bias.shape(), device)?.to_dtype(dtype)?;
122        // affine_mask: (B, L) → (B, 1, 1, L)
123        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        // Mask cross-chain pairs with -inf
130        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)?; // (B, 1, L, L)
134            let diff_bias = diff_chain.broadcast_as(attn_bias.shape())?;
135            attn_bias = diff_bias.where_cond(&neg_inf, &attn_bias)?;
136        }
137
138        // ── Project hidden states ──────────────────────────────────────────
139        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        // ── Rotation-only vectors: Q_rot, K_rot, V ────────────────────────
148        // Reshape: (B, L, (h*c)) → (B, L, h, 3)
149        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        // Rotate local-frame vectors to global frame: (B, L, h_rot, 3)
152        let vec_rot = Affine3D::apply_rot(&affine.rot, &vec_rot)?;
153
154        let query_rot = vec_rot.narrow(D::Minus2, 0, self.v_heads)?; // (B, L, H, 3)
155        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        // ── Full-affine (rot+trans) vectors: Q_dist, K_dist ───────────────
163        let vec_dist = vec_dist.reshape((b, l, self.v_heads * 2, 3))?;
164        let vec_dist = affine.apply(&vec_dist)?; // (B, L, H*2, 3) in global frame
165        let query_dist = vec_dist.narrow(D::Minus2, 0, self.v_heads)?; // (B, L, H, 3)
166        let key_dist = vec_dist.narrow(D::Minus2, self.v_heads, self.v_heads)?;
167
168        // ── Rearrange for attention computation ───────────────────────────
169        // (B, L, H, 3) → (B, H, L, 3)
170        let query_rot = query_rot.permute((0, 2, 1, 3))?;
171        // (B, L, H, 3) → (B, H, 3, L)  [for matmul with query]
172        let key_rot = key_rot.permute((0, 2, 3, 1))?;
173        // (B, L, H, 3) → (B, H, L, 1, 3)
174        let query_dist = query_dist.permute((0, 2, 1, 3))?.unsqueeze(D::Minus2)?;
175        // (B, L, H, 3) → (B, H, 1, L, 3)  [unsqueeze at position 2 = -3 of 5-dim result]
176        let key_dist = key_dist.permute((0, 2, 1, 3))?.unsqueeze(2)?;
177        // (B, L, H*num_vm, 3) → (B, H, L, num_vm*3)
178        let value = value
179            .reshape((b, l, self.v_heads, self.num_vector_messages * 3))?
180            .permute((0, 2, 1, 3))?;
181
182        // ── Attention weight: rotation + distance terms ────────────────────
183        // Rotation term: (B, H, L, 3) @ (B, H, 3, L) = (B, H, L, L)
184        // affine(scale, 0.0) = tensor * scale (scalar multiplication)
185        let rotation_term = query_rot
186            .contiguous()?
187            .matmul(&key_rot.contiguous()?)?
188            .affine(SQRT_3.recip(), 0.0)?;
189
190        // Distance term: ||q - k||_2 / sqrt(3) → (B, H, L, L)
191        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        // Learnable per-head weights: (H,) → (1, H, 1, 1); softplus = log(1 + exp(x))
201        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        // Add attention bias (already (B, 1, L, L); broadcast to (B, H, L, L))
213        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        // ── Weighted sum of values ────────────────────────────────────────
219        // (B, H, L, L) @ (B, H, L, num_vm*3) = (B, H, L, num_vm*3)
220        let attn_out = attn_weight.matmul(&value.contiguous()?)?;
221
222        // Rearrange: (B, H, L, num_vm*3) → (B, L, H*num_vm, 3)
223        let attn_out = attn_out
224            .permute((0, 2, 1, 3))? // (B, L, H, num_vm*3)
225            .contiguous()?
226            .reshape((b, l, self.v_heads * self.num_vector_messages, 3))?;
227
228        // Rotate back from global frame to local frame
229        let attn_out = Affine3D::apply_rot_inv(&affine.rot, &attn_out)?;
230
231        // Flatten head and vector-message dims: (B, L, H*num_vm, 3) → (B, L, H*num_vm*3)
232        let mut attn_out =
233            attn_out
234                .contiguous()?
235                .reshape((b, l, self.v_heads * self.num_vector_messages * 3))?;
236
237        // Zero out frameless positions if requested
238        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
250/// Numerically stable softplus: `log(1 + exp(x))`.
251fn softplus(t: &Tensor) -> Result<Tensor> {
252    // affine(1.0, 1.0) computes tensor * 1 + 1 = tensor + 1
253    t.exp()?.affine(1.0, 1.0)?.log()
254}