Skip to main content

ferritin_plms/
plm_runner.rs

1//! Unified trait for protein language model runners.
2//!
3//! `PlmRunner` provides a common interface for sequence embedding across
4//! ESM2, AMPLIFY, and ESMC. Downstream code can be generic over the runner
5//! type (e.g., for benchmarking or ensemble inference).
6
7use anyhow::Result;
8use candle_core::Tensor;
9
10/// Trait implemented by all PLM runner types.
11pub trait PlmRunner {
12    /// Run a forward pass on `sequence` and return per-residue embeddings.
13    ///
14    /// Shape: `(1, L, d_model)` where `L` includes any BOS/EOS tokens.
15    fn embed(&self, sequence: &str) -> Result<Tensor>;
16
17    /// Model name / identifier string (e.g. "esm2", "amplify", "esmc").
18    fn model_name(&self) -> &str;
19}
20
21#[cfg(test)]
22mod tests {
23    use super::*;
24    use candle_core::{DType, Device};
25
26    struct MockRunner;
27
28    impl PlmRunner for MockRunner {
29        fn embed(&self, _sequence: &str) -> Result<Tensor> {
30            Tensor::zeros((1usize, 3usize, 16usize), DType::F32, &Device::Cpu)
31                .map_err(anyhow::Error::from)
32        }
33
34        fn model_name(&self) -> &str {
35            "mock"
36        }
37    }
38
39    #[test]
40    fn test_mock_runner_model_name() {
41        let runner = MockRunner;
42        assert_eq!(runner.model_name(), "mock");
43    }
44
45    #[test]
46    fn test_mock_runner_embed_shape() {
47        let runner = MockRunner;
48        let tensor = runner.embed("ACDE").unwrap();
49        assert_eq!(tensor.dims(), &[1, 3, 16]);
50    }
51}