use ferrum_kernels::{backend::Backend, MarlinExpertStack};
use ferrum_types::{FerrumError, Result};
use crate::config::QuantConfig;
use crate::traits::Linear;
pub trait WeightLoader<B: Backend>: Send + Sync {
fn load_tensor(&self, name: &str) -> Result<B::Buffer>;
fn load_linear(&self, name: &str) -> Result<Box<dyn Linear<B>>>;
fn has_tensor(&self, name: &str) -> bool;
fn quant_config(&self) -> Option<&QuantConfig>;
fn load_stacked_gptq_experts(
&self,
expert_prefix_fmt: &str,
num_experts: usize,
proj_names: &[&str],
) -> Result<(std::sync::Arc<dyn MarlinExpertStack<B>>, usize, usize)> {
let _ = (expert_prefix_fmt, num_experts, proj_names);
Err(FerrumError::unsupported(
"load_stacked_gptq_experts not implemented for this weight loader",
))
}
}
pub struct PrefixedLoader<'a, B: Backend> {
inner: &'a dyn WeightLoader<B>,
prefix: String,
}
impl<'a, B: Backend> PrefixedLoader<'a, B> {
pub fn new(inner: &'a dyn WeightLoader<B>, prefix: impl Into<String>) -> Self {
Self {
inner,
prefix: prefix.into(),
}
}
}
impl<'a, B: Backend> WeightLoader<B> for PrefixedLoader<'a, B> {
fn load_tensor(&self, name: &str) -> Result<B::Buffer> {
self.inner.load_tensor(&format!("{}{}", self.prefix, name))
}
fn load_linear(&self, name: &str) -> Result<Box<dyn Linear<B>>> {
self.inner.load_linear(&format!("{}{}", self.prefix, name))
}
fn has_tensor(&self, name: &str) -> bool {
self.inner.has_tensor(&format!("{}{}", self.prefix, name))
}
fn quant_config(&self) -> Option<&QuantConfig> {
self.inner.quant_config()
}
fn load_stacked_gptq_experts(
&self,
expert_prefix_fmt: &str,
num_experts: usize,
proj_names: &[&str],
) -> Result<(std::sync::Arc<dyn MarlinExpertStack<B>>, usize, usize)> {
self.inner.load_stacked_gptq_experts(
&format!("{}{}", self.prefix, expert_prefix_fmt),
num_experts,
proj_names,
)
}
}