ferritin_plms/
plm_runner.rs1use anyhow::Result;
8use candle_core::Tensor;
9
10pub trait PlmRunner {
12 fn embed(&self, sequence: &str) -> Result<Tensor>;
16
17 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}