1use async_trait::async_trait;
14use ferrum_interfaces::{
15 KvCacheManager, ModelExecutor, Sampler, SchedulerInterface as Scheduler, Tokenizer,
16};
17use ferrum_models::vnext::{PreparedProductionModel, ProductionModelSourceBundle};
18use ferrum_types::{Device, EngineConfig, FerrumError, Result, RuntimeKnobs};
19use parking_lot::RwLock;
20use std::collections::HashMap;
21use std::sync::{Arc, OnceLock};
22use tracing::{debug, info};
23
24pub use ferrum_interfaces::sampler::GreedySampler;
25
26#[async_trait]
32pub trait ComponentFactory<T>: Send + Sync {
33 async fn create(&self, config: &ComponentConfig) -> Result<T>;
35
36 fn metadata(&self) -> ComponentMetadata;
38}
39
40#[derive(Debug, Clone)]
42pub struct ComponentConfig {
43 pub engine_config: EngineConfig,
45 pub device: Device,
47 pub component_options: HashMap<String, serde_json::Value>,
49 pub model_sources: Option<Arc<ProductionModelSourceBundle>>,
52 pub prepared_model: Option<Arc<PreparedProductionModel>>,
55}
56
57impl ComponentConfig {
58 pub fn from_engine_config(config: &EngineConfig) -> Self {
60 Self::from_engine_config_and_product_model(config, None, None)
61 }
62
63 pub fn from_engine_config_and_sources(
64 config: &EngineConfig,
65 model_sources: Option<Arc<ProductionModelSourceBundle>>,
66 ) -> Self {
67 Self::from_engine_config_and_product_model(config, model_sources, None)
68 }
69
70 pub fn from_engine_config_and_product_model(
71 config: &EngineConfig,
72 model_sources: Option<Arc<ProductionModelSourceBundle>>,
73 prepared_model: Option<Arc<PreparedProductionModel>>,
74 ) -> Self {
75 Self {
76 engine_config: config.clone(),
77 device: config.backend.device.clone(),
78 component_options: config.backend.backend_options.clone(),
79 model_sources,
80 prepared_model,
81 }
82 }
83
84 pub fn get_option<T: serde::de::DeserializeOwned>(&self, key: &str) -> Option<T> {
86 self.component_options
87 .get(key)
88 .and_then(|v| serde_json::from_value(v.clone()).ok())
89 }
90
91 pub fn get_string_option(&self, key: &str) -> Option<String> {
93 self.component_options
94 .get(key)
95 .and_then(|v| v.as_str())
96 .map(|s| s.to_string())
97 }
98}
99
100#[derive(Debug, Clone)]
102pub struct ComponentMetadata {
103 pub name: String,
105 pub version: String,
107 pub description: String,
109 pub supported_devices: Vec<Device>,
111 pub capabilities: Vec<String>,
113}
114
115impl Default for ComponentMetadata {
116 fn default() -> Self {
117 Self {
118 name: "unknown".to_string(),
119 version: "0.0.0".to_string(),
120 description: String::new(),
121 supported_devices: vec![Device::CPU],
122 capabilities: vec![],
123 }
124 }
125}
126
127fn cpu_cuda_and_optional_metal_devices() -> Vec<Device> {
128 #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
129 {
130 vec![Device::CPU, Device::CUDA(0), Device::Metal]
131 }
132 #[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
133 {
134 vec![Device::CPU, Device::CUDA(0)]
135 }
136}
137
138fn parse_executor_dtype(s: &str) -> Option<candle_core::DType> {
139 match s.trim().to_ascii_lowercase().as_str() {
140 "fp16" | "f16" | "float16" => Some(candle_core::DType::F16),
141 "fp32" | "f32" | "float32" => Some(candle_core::DType::F32),
142 _ => None,
143 }
144}
145
146#[derive(Debug, Clone, PartialEq, Eq)]
147struct RegistryRuntimeEnv {
148 model_path: Option<String>,
149 metal_dtype: Option<String>,
150 dtype: Option<String>,
151 tp: usize,
152}
153
154impl RegistryRuntimeEnv {
155 fn from_runtime_knobs(knobs: &RuntimeKnobs) -> Self {
160 Self {
161 model_path: knobs.model_path.clone(),
162 metal_dtype: knobs.metal_dtype.clone(),
163 dtype: knobs.dtype.clone(),
164 tp: knobs.tp.unwrap_or(0),
165 }
166 }
167
168 #[cfg(test)]
169 fn from_env_vars<I, K, V>(vars: I) -> Self
170 where
171 I: IntoIterator<Item = (K, V)>,
172 K: AsRef<str>,
173 V: Into<String>,
174 {
175 let mut model_path = None;
176 let mut metal_dtype = None;
177 let mut dtype = None;
178 let mut tp = None;
179
180 for (key, value) in vars {
181 let value = value.into();
182 match key.as_ref() {
183 "FERRUM_MODEL_PATH" => model_path = Some(value),
184 "FERRUM_METAL_DTYPE" => metal_dtype = Some(value),
185 "FERRUM_DTYPE" => dtype = Some(value),
186 "FERRUM_TP" => tp = value.parse::<usize>().ok(),
187 _ => {}
188 }
189 }
190
191 Self {
192 model_path,
193 metal_dtype,
194 dtype,
195 tp: tp.unwrap_or(0),
196 }
197 }
198
199 #[cfg(test)]
200 fn model_path(&self) -> Option<String> {
201 self.model_path.clone()
202 }
203
204 fn dtype_for_device(&self, device: &Device) -> candle_core::DType {
205 match device {
206 Device::CPU => candle_core::DType::F32,
207 #[cfg(any(target_os = "macos", target_os = "ios"))]
208 Device::Metal => self
209 .metal_dtype
210 .as_deref()
211 .and_then(parse_executor_dtype)
212 .or_else(|| self.dtype.as_deref().and_then(parse_executor_dtype))
213 .unwrap_or(candle_core::DType::F32),
214 Device::CUDA(_) | Device::ROCm(_) => self
215 .dtype
216 .as_deref()
217 .and_then(parse_executor_dtype)
218 .unwrap_or(candle_core::DType::F16),
219 }
220 }
221
222 #[allow(dead_code)]
223 fn metal_dtype_hint(&self) -> String {
224 self.metal_dtype
225 .clone()
226 .unwrap_or_else(|| "f32".to_string())
227 }
228}
229
230pub struct ComponentRegistry {
236 tokenizer_factories:
237 RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>>>,
238 sampler_factories:
239 RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>>>,
240 scheduler_factories:
241 RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>>>,
242 kv_cache_factories:
243 RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>>>,
244 executor_factories:
245 RwLock<HashMap<String, Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>>>,
246}
247
248impl ComponentRegistry {
249 pub fn new() -> Self {
251 Self {
252 tokenizer_factories: RwLock::new(HashMap::new()),
253 sampler_factories: RwLock::new(HashMap::new()),
254 scheduler_factories: RwLock::new(HashMap::new()),
255 kv_cache_factories: RwLock::new(HashMap::new()),
256 executor_factories: RwLock::new(HashMap::new()),
257 }
258 }
259
260 pub fn with_defaults() -> Self {
262 let registry = Self::new();
263 registry.register_defaults();
264 registry
265 }
266
267 pub fn register_defaults(&self) {
269 info!("Registering default component factories");
270
271 self.register_tokenizer_factory("huggingface", Arc::new(HuggingFaceTokenizerFactory));
273 self.register_tokenizer_factory("stub", Arc::new(StubTokenizerFactory));
274
275 self.register_sampler_factory("multinomial", Arc::new(MultinomialSamplerFactory));
277 self.register_sampler_factory("greedy", Arc::new(GreedySamplerFactory));
278
279 self.register_scheduler_factory("fifo", Arc::new(FifoSchedulerFactory));
281 self.register_scheduler_factory("priority", Arc::new(PrioritySchedulerFactory));
282 self.register_scheduler_factory("continuous", Arc::new(ContinuousBatchSchedulerFactory));
283
284 self.register_kv_cache_factory("default", Arc::new(DefaultKvCacheFactory));
286 self.register_kv_cache_factory("paged", Arc::new(PagedKvCacheFactory));
287
288 self.register_executor_factory("stub", Arc::new(StubExecutorFactory));
290 self.register_executor_factory("llm", Arc::new(LlmExecutorFactory));
291
292 debug!(
293 "Registered factories - tokenizers: {:?}, samplers: {:?}, schedulers: {:?}, kv_caches: {:?}, executors: {:?}",
294 self.list_tokenizers(),
295 self.list_samplers(),
296 self.list_schedulers(),
297 self.list_kv_caches(),
298 self.list_executors()
299 );
300 }
301
302 pub fn register_tokenizer_factory(
308 &self,
309 name: impl Into<String>,
310 factory: Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>,
311 ) {
312 let name = name.into();
313 debug!("Registering tokenizer factory: {}", name);
314 self.tokenizer_factories.write().insert(name, factory);
315 }
316
317 pub fn register_sampler_factory(
319 &self,
320 name: impl Into<String>,
321 factory: Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>,
322 ) {
323 let name = name.into();
324 debug!("Registering sampler factory: {}", name);
325 self.sampler_factories.write().insert(name, factory);
326 }
327
328 pub fn register_scheduler_factory(
330 &self,
331 name: impl Into<String>,
332 factory: Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>,
333 ) {
334 let name = name.into();
335 debug!("Registering scheduler factory: {}", name);
336 self.scheduler_factories.write().insert(name, factory);
337 }
338
339 pub fn register_kv_cache_factory(
341 &self,
342 name: impl Into<String>,
343 factory: Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>,
344 ) {
345 let name = name.into();
346 debug!("Registering KV cache factory: {}", name);
347 self.kv_cache_factories.write().insert(name, factory);
348 }
349
350 pub fn register_executor_factory(
352 &self,
353 name: impl Into<String>,
354 factory: Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>,
355 ) {
356 let name = name.into();
357 debug!("Registering executor factory: {}", name);
358 self.executor_factories.write().insert(name, factory);
359 }
360
361 pub fn get_tokenizer_factory(
367 &self,
368 name: &str,
369 ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Tokenizer + Send + Sync>>>> {
370 self.tokenizer_factories.read().get(name).cloned()
371 }
372
373 pub fn get_sampler_factory(
375 &self,
376 name: &str,
377 ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Sampler + Send + Sync>>>> {
378 self.sampler_factories.read().get(name).cloned()
379 }
380
381 pub fn get_scheduler_factory(
383 &self,
384 name: &str,
385 ) -> Option<Arc<dyn ComponentFactory<Arc<dyn Scheduler + Send + Sync>>>> {
386 self.scheduler_factories.read().get(name).cloned()
387 }
388
389 pub fn get_kv_cache_factory(
391 &self,
392 name: &str,
393 ) -> Option<Arc<dyn ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>>>> {
394 self.kv_cache_factories.read().get(name).cloned()
395 }
396
397 pub fn get_executor_factory(
399 &self,
400 name: &str,
401 ) -> Option<Arc<dyn ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>>>> {
402 self.executor_factories.read().get(name).cloned()
403 }
404
405 pub fn list_tokenizers(&self) -> Vec<String> {
411 self.tokenizer_factories.read().keys().cloned().collect()
412 }
413
414 pub fn list_samplers(&self) -> Vec<String> {
416 self.sampler_factories.read().keys().cloned().collect()
417 }
418
419 pub fn list_schedulers(&self) -> Vec<String> {
421 self.scheduler_factories.read().keys().cloned().collect()
422 }
423
424 pub fn list_kv_caches(&self) -> Vec<String> {
426 self.kv_cache_factories.read().keys().cloned().collect()
427 }
428
429 pub fn list_executors(&self) -> Vec<String> {
431 self.executor_factories.read().keys().cloned().collect()
432 }
433
434 pub async fn create_tokenizer(
440 &self,
441 name: &str,
442 config: &ComponentConfig,
443 ) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
444 let factory = self.get_tokenizer_factory(name).ok_or_else(|| {
445 FerrumError::tokenizer(format!(
446 "Tokenizer '{}' not found. Available: {:?}",
447 name,
448 self.list_tokenizers()
449 ))
450 })?;
451 factory.create(config).await
452 }
453
454 pub async fn create_sampler(
456 &self,
457 name: &str,
458 config: &ComponentConfig,
459 ) -> Result<Arc<dyn Sampler + Send + Sync>> {
460 let factory = self.get_sampler_factory(name).ok_or_else(|| {
461 FerrumError::internal(format!(
462 "Sampler '{}' not found. Available: {:?}",
463 name,
464 self.list_samplers()
465 ))
466 })?;
467 factory.create(config).await
468 }
469
470 pub async fn create_scheduler(
472 &self,
473 name: &str,
474 config: &ComponentConfig,
475 ) -> Result<Arc<dyn Scheduler + Send + Sync>> {
476 let factory = self.get_scheduler_factory(name).ok_or_else(|| {
477 FerrumError::scheduler(format!(
478 "Scheduler '{}' not found. Available: {:?}",
479 name,
480 self.list_schedulers()
481 ))
482 })?;
483 factory.create(config).await
484 }
485
486 pub async fn create_kv_cache(
488 &self,
489 name: &str,
490 config: &ComponentConfig,
491 ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
492 let factory = self.get_kv_cache_factory(name).ok_or_else(|| {
493 FerrumError::internal(format!(
494 "KV cache '{}' not found. Available: {:?}",
495 name,
496 self.list_kv_caches()
497 ))
498 })?;
499 factory.create(config).await
500 }
501
502 pub async fn create_executor(
504 &self,
505 name: &str,
506 config: &ComponentConfig,
507 ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
508 let factory = self.get_executor_factory(name).ok_or_else(|| {
509 FerrumError::model(format!(
510 "Executor '{}' not found. Available: {:?}",
511 name,
512 self.list_executors()
513 ))
514 })?;
515 factory.create(config).await
516 }
517}
518
519impl Default for ComponentRegistry {
520 fn default() -> Self {
521 Self::with_defaults()
522 }
523}
524
525impl std::fmt::Debug for ComponentRegistry {
526 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
527 f.debug_struct("ComponentRegistry")
528 .field("tokenizers", &self.list_tokenizers())
529 .field("samplers", &self.list_samplers())
530 .field("schedulers", &self.list_schedulers())
531 .field("kv_caches", &self.list_kv_caches())
532 .field("executors", &self.list_executors())
533 .finish()
534 }
535}
536
537pub struct HuggingFaceTokenizerFactory;
557
558#[async_trait]
559impl ComponentFactory<Arc<dyn Tokenizer + Send + Sync>> for HuggingFaceTokenizerFactory {
560 async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
561 if let Some(sources) = config.model_sources.as_ref() {
562 let tokenizer =
563 ferrum_tokenizer::implementations::HuggingFaceTokenizer::from_source_bytes(
564 sources.tokenizer_json(),
565 sources.tokenizer_config_json(),
566 sources.generation_config_json(),
567 )
568 .await?;
569 info!(
570 tokenizer_root = %sources.tokenizer_root().display(),
571 "Loaded HuggingFace tokenizer from typed product sources"
572 );
573 return Ok(Arc::new(tokenizer));
574 }
575
576 let tokenizer_path = config
578 .get_string_option("tokenizer_path")
579 .or_else(|| config.get_string_option("model_path"))
580 .or_else(|| config.engine_config.runtime.model_path.clone());
581
582 if let Some(model_path) = tokenizer_path {
583 let path = std::path::Path::new(&model_path);
584 let tokenizer_file = if ferrum_models::gguf_engine_loader::is_gguf_path(&model_path) {
587 ferrum_models::gguf_engine_loader::auto_discover_tokenizer_path(path).ok_or_else(
588 || {
589 FerrumError::tokenizer(format!(
590 "Could not find tokenizer.json for {} — \
591 place it next to the .gguf file or in a \
592 sibling tokenizers/ directory",
593 path.display()
594 ))
595 },
596 )?
597 } else {
598 path.join("tokenizer.json")
600 };
601
602 if tokenizer_file.exists() {
603 info!("Loading HuggingFace tokenizer from: {:?}", tokenizer_file);
604 match ferrum_tokenizer::implementations::HuggingFaceTokenizer::from_file(
605 &tokenizer_file.to_string_lossy(),
606 )
607 .await
608 {
609 Ok(tokenizer) => {
610 info!("HuggingFace tokenizer loaded successfully");
611 return Ok(Arc::new(tokenizer));
612 }
613 Err(e) => {
614 tracing::warn!("Failed to load tokenizer: {}, falling back to stub", e);
615 }
616 }
617 }
618 }
619
620 Err(FerrumError::tokenizer(
622 "HuggingFace tokenizer path not found or invalid",
623 ))
624 }
625
626 fn metadata(&self) -> ComponentMetadata {
627 ComponentMetadata {
628 name: "huggingface".to_string(),
629 version: "0.1.0".to_string(),
630 description: "HuggingFace tokenizers library integration".to_string(),
631 supported_devices: vec![Device::CPU],
632 capabilities: vec![
633 "bpe".to_string(),
634 "wordpiece".to_string(),
635 "sentencepiece".to_string(),
636 "chat_template".to_string(),
637 ],
638 }
639 }
640}
641
642pub struct StubTokenizerFactory;
644
645#[async_trait]
646impl ComponentFactory<Arc<dyn Tokenizer + Send + Sync>> for StubTokenizerFactory {
647 async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Tokenizer + Send + Sync>> {
648 let vocab_size = config
649 .engine_config
650 .model
651 .model_info
652 .as_ref()
653 .map(|info| info.vocab_size)
654 .unwrap_or(32000);
655
656 info!("Creating stub tokenizer with vocab_size: {}", vocab_size);
657 Ok(Arc::new(StubTokenizer::new(vocab_size)))
658 }
659
660 fn metadata(&self) -> ComponentMetadata {
661 ComponentMetadata {
662 name: "stub".to_string(),
663 version: "0.1.0".to_string(),
664 description: "Stub tokenizer for testing".to_string(),
665 supported_devices: vec![Device::CPU],
666 capabilities: vec!["testing".to_string()],
667 }
668 }
669}
670
671pub struct MultinomialSamplerFactory;
677
678#[async_trait]
679impl ComponentFactory<Arc<dyn Sampler + Send + Sync>> for MultinomialSamplerFactory {
680 async fn create(&self, _config: &ComponentConfig) -> Result<Arc<dyn Sampler + Send + Sync>> {
681 info!("Creating multinomial sampler");
682 Ok(Arc::new(ferrum_interfaces::sampler::MultinomialSampler))
683 }
684
685 fn metadata(&self) -> ComponentMetadata {
686 ComponentMetadata {
687 name: "multinomial".to_string(),
688 version: "0.1.0".to_string(),
689 description: "Multinomial sampling with temperature and top-k/top-p".to_string(),
690 supported_devices: vec![Device::CPU],
691 capabilities: vec![
692 "temperature".to_string(),
693 "top_k".to_string(),
694 "top_p".to_string(),
695 ],
696 }
697 }
698}
699
700pub struct GreedySamplerFactory;
702
703#[async_trait]
704impl ComponentFactory<Arc<dyn Sampler + Send + Sync>> for GreedySamplerFactory {
705 async fn create(&self, _config: &ComponentConfig) -> Result<Arc<dyn Sampler + Send + Sync>> {
706 info!("Creating greedy sampler");
707 Ok(Arc::new(GreedySampler))
708 }
709
710 fn metadata(&self) -> ComponentMetadata {
711 ComponentMetadata {
712 name: "greedy".to_string(),
713 version: "0.1.0".to_string(),
714 description: "Greedy decoding (always pick highest probability)".to_string(),
715 supported_devices: vec![Device::CPU],
716 capabilities: vec!["deterministic".to_string()],
717 }
718 }
719}
720
721pub struct FifoSchedulerFactory;
727
728#[async_trait]
729impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for FifoSchedulerFactory {
730 async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
731 info!("Creating FIFO scheduler");
732 let scheduler_config = config.engine_config.scheduler.clone();
733 let scheduler = ferrum_scheduler::implementations::FifoScheduler::new(scheduler_config);
734 Ok(Arc::new(scheduler))
735 }
736
737 fn metadata(&self) -> ComponentMetadata {
738 ComponentMetadata {
739 name: "fifo".to_string(),
740 version: "0.1.0".to_string(),
741 description: "First-In-First-Out scheduler".to_string(),
742 supported_devices: vec![Device::CPU],
743 capabilities: vec!["simple".to_string(), "fair".to_string()],
744 }
745 }
746}
747
748pub struct PrioritySchedulerFactory;
750
751#[async_trait]
752impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for PrioritySchedulerFactory {
753 async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
754 info!("Creating priority scheduler");
755 let scheduler_config = config.engine_config.scheduler.clone();
756 let scheduler = ferrum_scheduler::implementations::PriorityScheduler::new(scheduler_config);
757 Ok(Arc::new(scheduler))
758 }
759
760 fn metadata(&self) -> ComponentMetadata {
761 ComponentMetadata {
762 name: "priority".to_string(),
763 version: "0.1.0".to_string(),
764 description: "Priority-based scheduler".to_string(),
765 supported_devices: vec![Device::CPU],
766 capabilities: vec!["priority".to_string(), "preemption".to_string()],
767 }
768 }
769}
770
771pub struct ContinuousBatchSchedulerFactory;
773
774#[async_trait]
775impl ComponentFactory<Arc<dyn Scheduler + Send + Sync>> for ContinuousBatchSchedulerFactory {
776 async fn create(&self, config: &ComponentConfig) -> Result<Arc<dyn Scheduler + Send + Sync>> {
777 info!("Creating continuous batch scheduler");
778 let scheduler_config = config.engine_config.scheduler.clone();
779 let scheduler =
780 ferrum_scheduler::implementations::ContinuousBatchScheduler::new(scheduler_config);
781 Ok(Arc::new(scheduler))
782 }
783
784 fn metadata(&self) -> ComponentMetadata {
785 ComponentMetadata {
786 name: "continuous".to_string(),
787 version: "0.1.0".to_string(),
788 description: "Continuous batching scheduler with iteration-level scheduling"
789 .to_string(),
790 supported_devices: cpu_cuda_and_optional_metal_devices(),
791 capabilities: vec![
792 "continuous_batching".to_string(),
793 "preemption".to_string(),
794 "chunked_prefill".to_string(),
795 "iteration_level".to_string(),
796 ],
797 }
798 }
799}
800
801pub struct DefaultKvCacheFactory;
807
808#[async_trait]
809impl ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>> for DefaultKvCacheFactory {
810 async fn create(
811 &self,
812 config: &ComponentConfig,
813 ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
814 let block_size = config.engine_config.kv_cache.block_size;
815 let max_blocks = config.engine_config.kv_cache.max_blocks;
816
817 info!(
818 "Creating default KV cache manager: device={:?}, block_size={}, max_blocks={}",
819 config.device, block_size, max_blocks
820 );
821
822 let manager = ferrum_kv::managers::DefaultKvCacheManager::new(
823 config.device.clone(),
824 block_size,
825 max_blocks,
826 )?;
827 Ok(Arc::new(manager))
828 }
829
830 fn metadata(&self) -> ComponentMetadata {
831 ComponentMetadata {
832 name: "default".to_string(),
833 version: "0.1.0".to_string(),
834 description: "Default contiguous KV cache manager".to_string(),
835 supported_devices: cpu_cuda_and_optional_metal_devices(),
836 capabilities: vec!["contiguous".to_string()],
837 }
838 }
839}
840
841pub struct PagedKvCacheFactory;
843
844#[async_trait]
845impl ComponentFactory<Arc<dyn KvCacheManager + Send + Sync>> for PagedKvCacheFactory {
846 async fn create(
847 &self,
848 config: &ComponentConfig,
849 ) -> Result<Arc<dyn KvCacheManager + Send + Sync>> {
850 let block_size = config.engine_config.kv_cache.block_size;
851 let max_blocks = config.engine_config.kv_cache.max_blocks;
852
853 info!(
854 "Creating paged KV cache manager: device={:?}, block_size={}, max_blocks={}",
855 config.device, block_size, max_blocks
856 );
857
858 let paged_config = ferrum_kv::managers::PagedKvCacheConfig {
860 block_size,
861 max_gpu_blocks: max_blocks,
862 max_cpu_blocks: max_blocks / 2,
863 enable_cow: true,
864 enable_swapping: true,
865 ..Default::default()
866 };
867
868 let manager =
869 ferrum_kv::managers::PagedKvCacheManager::new(config.device.clone(), paged_config)?;
870 Ok(Arc::new(manager))
871 }
872
873 fn metadata(&self) -> ComponentMetadata {
874 ComponentMetadata {
875 name: "paged".to_string(),
876 version: "0.1.0".to_string(),
877 description: "Paged KV cache manager for PagedAttention".to_string(),
878 supported_devices: cpu_cuda_and_optional_metal_devices(),
879 capabilities: vec![
880 "paged".to_string(),
881 "copy_on_write".to_string(),
882 "swap".to_string(),
883 ],
884 }
885 }
886}
887
888pub struct StubExecutorFactory;
894
895#[async_trait]
896impl ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>> for StubExecutorFactory {
897 async fn create(
898 &self,
899 config: &ComponentConfig,
900 ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
901 let vocab_size = config
902 .engine_config
903 .model
904 .model_info
905 .as_ref()
906 .map(|info| info.vocab_size)
907 .unwrap_or(32000);
908
909 info!(
910 "Creating stub executor for model: {}",
911 config.engine_config.model.model_id
912 );
913
914 let tensor_factory: Arc<dyn ferrum_interfaces::TensorFactory> = Arc::new(
919 crate::tensor_factory::candle::CandleTensorFactory::new(config.device.clone()),
920 );
921
922 let executor = ferrum_models::StubModelExecutor::new(
923 config.engine_config.model.model_id.clone(),
924 vocab_size,
925 tensor_factory,
926 );
927
928 Ok(Arc::new(executor))
929 }
930
931 fn metadata(&self) -> ComponentMetadata {
932 ComponentMetadata {
933 name: "stub".to_string(),
934 version: "0.1.0".to_string(),
935 description: "Stub executor for testing".to_string(),
936 supported_devices: vec![Device::CPU],
937 capabilities: vec!["testing".to_string()],
938 }
939 }
940}
941
942pub struct LlmExecutorFactory;
960
961fn resolve_llama_layer_split_plan(
962 config: &ComponentConfig,
963 num_layers: usize,
964) -> Result<Option<crate::layer_split::ParsedLayerSplitPlan>> {
965 if config
966 .get_string_option("selected_distributed_strategy")
967 .as_deref()
968 != Some("layer_split")
969 {
970 return Ok(None);
971 }
972
973 let selected = config
974 .get_option::<Vec<usize>>("selected_gpu_devices")
975 .unwrap_or_default();
976 let plan_raw = config.get_string_option("selected_layer_split_plan");
977 let parsed_plan =
978 if let Some(stages) = config.component_options.get("selected_layer_split_stages") {
979 crate::layer_split::parse_layer_split_stage_documents(stages)?
980 } else {
981 let plan_raw = plan_raw.as_deref().ok_or_else(|| {
982 FerrumError::config(
983 "selected_distributed_strategy=layer_split requires selected_layer_split_plan",
984 )
985 })?;
986 crate::layer_split::parse_layer_split_plan(plan_raw)?
987 };
988 crate::layer_split::validate_layer_split_plan_for_devices(&parsed_plan, &selected)?;
989 if parsed_plan.total_layers() != num_layers {
990 return Err(FerrumError::config(format!(
991 "selected_layer_split_plan covers {} layers but model has {num_layers}",
992 parsed_plan.total_layers()
993 )));
994 }
995 Ok(Some(parsed_plan))
996}
997
998fn resolve_llama_layer_split_pipeline_mode(
999 config: &ComponentConfig,
1000 stage_count: usize,
1001) -> Result<ferrum_models::models::LlamaPipelineMode> {
1002 let Some(mode) = config.get_string_option("layer_split_pipeline_mode") else {
1003 return Ok(ferrum_models::models::LlamaPipelineMode::default_for_stage_count(stage_count));
1004 };
1005 let mode = ferrum_models::models::LlamaPipelineMode::from_config_value(&mode)?;
1006 if mode == ferrum_models::models::LlamaPipelineMode::Overlapped && stage_count != 2 {
1007 return Err(FerrumError::config(
1008 "layer_split_pipeline_mode=overlapped requires exactly two pipeline stages",
1009 ));
1010 }
1011 Ok(mode)
1012}
1013
1014#[cfg(test)]
1015fn resolve_llama_layer_stage_config(
1016 config: &ComponentConfig,
1017 num_layers: usize,
1018) -> Result<Option<ferrum_models::models::llama_family::LlamaFamilyLayerStageConfig>> {
1019 let Some(parsed_plan) = resolve_llama_layer_split_plan(config, num_layers)? else {
1020 return Ok(None);
1021 };
1022 let device_id = match &config.device {
1023 Device::CUDA(device_id) => *device_id,
1024 other => {
1025 return Err(FerrumError::unsupported(format!(
1026 "selected_distributed_strategy=layer_split requires a CUDA stage device, got {other:?}",
1027 )));
1028 }
1029 };
1030 parsed_plan
1031 .llama_stage_config_for_device(device_id)
1032 .map(Some)
1033}
1034
1035fn build_llm<B, K>(
1045 arch: ferrum_models::Architecture,
1046 qcfg: ferrum_models::models::LlamaFamilyConfig,
1047 moe_cfg: Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1048 model_path: &str,
1049 llama_layer_split_plan: Option<crate::layer_split::ParsedLayerSplitPlan>,
1050 llama_layer_split_pipeline_mode: Option<ferrum_models::models::LlamaPipelineMode>,
1051) -> Result<Box<dyn ferrum_models::common::DecoderOnlyLLM>>
1052where
1053 B: ferrum_kernels::backend::MoeLlmBackend,
1054 K: ferrum_kernels::backend::KvLayer<B>,
1055 ferrum_models::models::LlamaFamilyModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1056 ferrum_models::models::LlamaFamilyPipelineModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1057{
1058 if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1059 if llama_layer_split_plan.is_some() {
1060 return Err(FerrumError::unsupported(
1061 "CUDA layer_split stage loading is wired only for Llama-family dense models; \
1062 Qwen3MoeModel requires a separate MoE stage loader.",
1063 ));
1064 }
1065 let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1066 let mc = moe_cfg.ok_or_else(|| {
1067 FerrumError::internal(
1068 "Qwen3Moe arch reached build_llm without Qwen3MoeConfig (caller bug)",
1069 )
1070 })?;
1071 Ok(Box::new(
1072 ferrum_models::models::Qwen3MoeModel::<B, K>::new_safetensors(mc, &weight_loader)?,
1073 ))
1074 } else {
1075 if llama_layer_split_plan.is_some() && !B::supports_device_ordinal_scope() {
1076 return Err(FerrumError::unsupported(
1077 "selected_distributed_strategy=layer_split requires a backend with \
1078 device-scoped execution; refusing to silently use the default device",
1079 ));
1080 }
1081 let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1082 if let Some(plan) = llama_layer_split_plan {
1083 let stage_configs = plan.to_llama_stage_configs();
1084 let stage_device_ordinals = plan
1085 .stages
1086 .iter()
1087 .map(|stage| Some(stage.device))
1088 .collect::<Vec<_>>();
1089 let mut stages = Vec::with_capacity(stage_configs.len());
1090 for (idx, (stage_config, device_ordinal)) in stage_configs
1091 .into_iter()
1092 .zip(stage_device_ordinals.iter().copied())
1093 .enumerate()
1094 {
1095 tracing::info!(
1096 "Loading Llama layer_split stage {idx} on backend device {:?}",
1097 device_ordinal
1098 );
1099 let stage = B::with_device_ordinal(device_ordinal, || {
1100 ferrum_models::models::LlamaFamilyModel::<B, K>::new_layer_stage(
1101 qcfg.clone(),
1102 &weight_loader,
1103 stage_config,
1104 )
1105 })?;
1106 stages.push(stage);
1107 }
1108 Ok(Box::new(ferrum_models::models::LlamaFamilyPipelineModel::<
1109 B,
1110 K,
1111 >::new_with_placement(
1112 stages,
1113 ferrum_models::models::LlamaPipelinePlacement::from_backend_device_ordinals(
1114 stage_device_ordinals,
1115 )
1116 .with_pipeline_mode(
1117 llama_layer_split_pipeline_mode.expect("layer split pipeline mode resolved"),
1118 ),
1119 )?))
1120 } else {
1121 Ok(Box::new(
1122 ferrum_models::models::LlamaFamilyModel::<B, K>::new(qcfg, &weight_loader)?,
1123 ))
1124 }
1125 }
1126}
1127
1128fn validate_registered_vnext_backend(
1129 kind: ferrum_models::vnext::ProductionExecutionKind,
1130 device: &Device,
1131 external_metadata_id: &ferrum_interfaces::vnext::ExternalModelMetadataId,
1132) -> Result<()> {
1133 use ferrum_models::vnext::ProductionExecutionKind;
1134
1135 match (kind, device) {
1136 (ProductionExecutionKind::CausalLanguage, Device::CUDA(_)) => {
1137 #[cfg(feature = "cuda")]
1138 return Ok(());
1139 #[cfg(not(feature = "cuda"))]
1140 Err(FerrumError::device(
1141 "registered vNext CUDA composition requires the 'cuda' feature",
1142 ))
1143 }
1144 #[cfg(any(target_os = "macos", target_os = "ios"))]
1145 (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1146 #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1147 return Ok(());
1148 #[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
1149 Err(FerrumError::device(
1150 "registered vNext Metal composition requires the 'metal' feature on an Apple platform",
1151 ))
1152 }
1153 (kind, device) => Err(FerrumError::unsupported(format!(
1154 "registered vNext model metadata {external_metadata_id} requires a {kind:?} backend composition, but {device} is not registered"
1155 ))),
1156 }
1157}
1158
1159fn create_registered_vnext_executor(
1160 config: &ComponentConfig,
1161 model_path: &std::path::Path,
1162 sources: Option<Arc<ProductionModelSourceBundle>>,
1163 prepared_model: Option<Arc<PreparedProductionModel>>,
1164 registration: ferrum_models::vnext::RegisteredProductionModel,
1165) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1166 use ferrum_models::vnext::ProductionExecutionKind;
1167
1168 validate_registered_vnext_backend(
1169 registration.execution_kind(),
1170 &config.device,
1171 registration.external_metadata_id(),
1172 )?;
1173
1174 let _prepared_model_reused = prepared_model.is_some();
1175 if let (Some(prepared), Some(sources)) = (prepared_model.as_ref(), sources.as_ref()) {
1176 if !Arc::ptr_eq(prepared.sources(), sources) {
1177 return Err(FerrumError::model(
1178 "prepared product model and component sources do not share one source lease",
1179 ));
1180 }
1181 }
1182 let prepared = match prepared_model {
1183 Some(prepared) => {
1184 if prepared.family().external_metadata_id() != registration.external_metadata_id() {
1185 return Err(FerrumError::model(format!(
1186 "prepared product metadata {} differs from registered metadata {}",
1187 prepared.family().external_metadata_id(),
1188 registration.external_metadata_id()
1189 )));
1190 }
1191 prepared
1192 }
1193 None => Arc::new(match sources {
1194 Some(sources) => registration.prepare_from_sources(sources)?,
1195 None => registration.prepare(model_path)?,
1196 }),
1197 };
1198
1199 match (registration.execution_kind(), &config.device) {
1200 (ProductionExecutionKind::CausalLanguage, Device::CUDA(ordinal)) => {
1201 #[cfg(feature = "cuda")]
1202 {
1203 let model_info = prepared.model_info(
1204 config.engine_config.model.model_id.clone(),
1205 config.device.clone(),
1206 );
1207 let family = prepared.family();
1208 let family_fingerprint = family
1209 .fingerprint()
1210 .map_err(|error| FerrumError::model(error.to_string()))?;
1211 let program_fingerprint = family
1212 .program()
1213 .fingerprint()
1214 .map_err(|error| FerrumError::model(error.to_string()))?;
1215 info!(
1216 external_metadata_id = %registration.external_metadata_id(),
1217 family_id = %family.family_id(),
1218 family_fingerprint,
1219 program_fingerprint,
1220 prepared_model_reused = _prepared_model_reused,
1221 backend = "cuda",
1222 "Building registered model from a typed vNext execution plan"
1223 );
1224 let device_id = ferrum_interfaces::vnext::DeviceId::new(format!(
1225 "device.cuda.{ordinal}"
1226 ))
1227 .map_err(|error| FerrumError::device(error.to_string()))?;
1228 let composition =
1229 ferrum_kernels::backend::cuda::vnext_ops::CudaVNextComposition::create(
1230 *ordinal,
1231 device_id,
1232 config.engine_config.runtime.attention_execution_policy,
1233 )
1234 .map_err(|error| {
1235 FerrumError::device(format!("create vNext CUDA runtime: {error}"))
1236 })?;
1237 let (
1238 runtime,
1239 operation_registry,
1240 weight_materializers,
1241 weight_materializer_id,
1242 catalog,
1243 ) = composition.into_parts();
1244 let executor = crate::product_composition::create_vnext_executor(
1245 &config.engine_config,
1246 prepared.as_ref(),
1247 model_info,
1248 runtime,
1249 operation_registry,
1250 weight_materializers,
1251 weight_materializer_id,
1252 catalog,
1253 )?;
1254 info!(
1255 resolved_plan_fingerprint = executor
1256 .resolved_model_plan()
1257 .map(|plan| plan.fingerprint())
1258 .unwrap_or("missing"),
1259 "Resolved product model plan is authoritative for vNext execution"
1260 );
1261 Ok(Arc::new(executor))
1262 }
1263 #[cfg(not(feature = "cuda"))]
1264 {
1265 let _ = (ordinal, model_path, prepared);
1266 Err(FerrumError::device(
1267 "registered vNext CUDA composition requires the 'cuda' feature",
1268 ))
1269 }
1270 }
1271 #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1272 (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1273 let model_info = prepared.model_info(
1274 config.engine_config.model.model_id.clone(),
1275 config.device.clone(),
1276 );
1277 let family = prepared.family();
1278 let family_fingerprint = family
1279 .fingerprint()
1280 .map_err(|error| FerrumError::model(error.to_string()))?;
1281 let program_fingerprint = family
1282 .program()
1283 .fingerprint()
1284 .map_err(|error| FerrumError::model(error.to_string()))?;
1285 info!(
1286 external_metadata_id = %registration.external_metadata_id(),
1287 family_id = %family.family_id(),
1288 family_fingerprint,
1289 program_fingerprint,
1290 prepared_model_reused = _prepared_model_reused,
1291 backend = "metal",
1292 "Building registered model from a typed vNext execution plan"
1293 );
1294 let device_id = ferrum_interfaces::vnext::DeviceId::new("device.metal.0")
1295 .map_err(|error| FerrumError::device(error.to_string()))?;
1296 let composition =
1297 ferrum_kernels::backend::metal::vnext_ops::MetalVNextComposition::create(
1298 device_id,
1299 )
1300 .map_err(|error| {
1301 FerrumError::device(format!("create vNext Metal runtime: {error}"))
1302 })?;
1303 let (
1304 runtime,
1305 operation_registry,
1306 weight_materializers,
1307 weight_materializer_id,
1308 catalog,
1309 ) = composition.into_parts();
1310 let executor = crate::product_composition::create_vnext_executor(
1311 &config.engine_config,
1312 prepared.as_ref(),
1313 model_info,
1314 runtime,
1315 operation_registry,
1316 weight_materializers,
1317 weight_materializer_id,
1318 catalog,
1319 )?;
1320 info!(
1321 resolved_plan_fingerprint = executor
1322 .resolved_model_plan()
1323 .map(|plan| plan.fingerprint())
1324 .unwrap_or("missing"),
1325 "Resolved product model plan is authoritative for vNext execution"
1326 );
1327 Ok(Arc::new(executor))
1328 }
1329 (kind, device) => Err(FerrumError::unsupported(format!(
1330 "registered vNext model metadata {} requires a {kind:?} backend composition, but {device} is not registered",
1331 registration.external_metadata_id()
1332 ))),
1333 }
1334}
1335
1336#[deprecated(note = "use `LlmExecutorFactory` (renamed PR A — Dim 1/3 cleanup)")]
1339pub type CandleExecutorFactory = LlmExecutorFactory;
1340
1341#[async_trait]
1342impl ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>> for LlmExecutorFactory {
1343 async fn create(
1344 &self,
1345 config: &ComponentConfig,
1346 ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1347 use candle_core::{DType, Device as CandleDevice};
1348 use ferrum_models::weight_format::WeightFormat;
1349
1350 let checkpoint_capture_enabled = config
1352 .engine_config
1353 .runtime
1354 .vnext_checkpoint_capture
1355 .is_some();
1356 let model_sources = config.model_sources.clone();
1357 let prepared_model = config.prepared_model.clone();
1358 let model_path = model_sources
1359 .as_ref()
1360 .map(|sources| sources.weights().path().display().to_string())
1361 .or_else(|| config.get_string_option("model_path"))
1362 .or_else(|| config.engine_config.runtime.model_path.clone());
1363
1364 let model_path = match model_path {
1365 Some(path) => path,
1366 None => {
1367 if checkpoint_capture_enabled {
1368 return Err(FerrumError::unsupported(
1369 "vNext checkpoint capture requires a registered model source",
1370 ));
1371 }
1372 info!("No model path found, falling back to stub executor");
1373 return StubExecutorFactory.create(config).await;
1374 }
1375 };
1376
1377 let weight_fmt = WeightFormat::detect(std::path::Path::new(&model_path))?;
1382 info!(
1383 "Loading model from {} (format: {})",
1384 model_path,
1385 weight_fmt.label()
1386 );
1387
1388 let production_registration = match model_sources.as_ref() {
1397 Some(sources) => Some(ferrum_models::vnext::resolve_registered_model_from_sources(
1398 sources,
1399 )?),
1400 None if !matches!(weight_fmt, WeightFormat::Gguf { .. }) => {
1401 Some(ferrum_models::vnext::resolve_registered_model_from_dir(
1402 std::path::Path::new(&model_path),
1403 )?)
1404 }
1405 None => None,
1406 };
1407 if let Some(production_registration) = production_registration {
1408 match production_registration {
1409 ferrum_models::vnext::ProductionModelRegistration::Registered(registration) => {
1410 return create_registered_vnext_executor(
1411 config,
1412 std::path::Path::new(&model_path),
1413 model_sources,
1414 prepared_model,
1415 registration,
1416 );
1417 }
1418 ferrum_models::vnext::ProductionModelRegistration::LegacyRegistered {
1419 external_metadata_id,
1420 } => {
1421 if checkpoint_capture_enabled {
1422 return Err(FerrumError::unsupported(format!(
1423 "vNext checkpoint capture requires a migrated model; metadata {external_metadata_id} is still legacy"
1424 )));
1425 }
1426 info!(
1427 %external_metadata_id,
1428 "Entering the explicitly registered legacy model path"
1429 );
1430 }
1431 }
1432 }
1433 if checkpoint_capture_enabled {
1434 return Err(FerrumError::unsupported(
1435 "vNext checkpoint capture requires a registered vNext model package",
1436 ));
1437 }
1438
1439 if let WeightFormat::Gguf { ref path } = weight_fmt {
1440 let (llm, model_info) = ferrum_models::gguf_engine_loader::load_gguf_decoder_with_info(
1443 path,
1444 &config.device,
1445 config.engine_config.model.model_id.clone(),
1446 )?;
1447 return Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)));
1448 }
1449
1450 let mut config_manager = ferrum_models::ConfigManager::new();
1454 let model_def = config_manager
1455 .load_from_path(std::path::Path::new(&model_path))
1456 .await?;
1457
1458 info!(
1459 " Architecture: {:?}, Layers: {}, Vocab: {}",
1460 model_def.architecture, model_def.num_hidden_layers, model_def.vocab_size
1461 );
1462
1463 let legacy_candle_device = || -> Result<CandleDevice> {
1470 match &config.device {
1471 Device::CPU => Ok(CandleDevice::Cpu),
1472 #[cfg(feature = "candle-cuda-compat")]
1473 Device::CUDA(id) => CandleDevice::new_cuda(*id)
1474 .map_err(|e| FerrumError::device(format!("CUDA error: {}", e))),
1475 #[cfg(not(feature = "candle-cuda-compat"))]
1476 Device::CUDA(_) => Err(FerrumError::unsupported(
1477 "legacy Candle CUDA executors require the candle-cuda-compat feature",
1478 )),
1479 #[cfg(any(target_os = "macos", target_os = "ios"))]
1480 Device::Metal => CandleDevice::new_metal(0)
1481 .map_err(|e| FerrumError::device(format!("Metal error: {}", e))),
1482 Device::ROCm(_) => Err(FerrumError::device("ROCm not yet supported")),
1483 }
1484 };
1485
1486 let dtype: DType = RegistryRuntimeEnv::from_runtime_knobs(&config.engine_config.runtime)
1495 .dtype_for_device(&config.device);
1496
1497 info!("Building model...");
1499 match model_def.architecture {
1500 arch @ (ferrum_models::Architecture::Llama
1504 | ferrum_models::Architecture::Qwen2
1505 | ferrum_models::Architecture::Qwen3
1506 | ferrum_models::Architecture::Qwen3Moe
1507 | ferrum_models::Architecture::Gemma3
1508 | ferrum_models::Architecture::Mistral) => {
1509 let _loader = ferrum_models::SafeTensorsLoader::new(&model_path);
1510 let model_dir_path: std::path::PathBuf = model_path.clone().into();
1511
1512 let tp_size = config.engine_config.runtime.tp.unwrap_or(0);
1517 if tp_size > 1 {
1518 return Err(FerrumError::unsupported(
1519 "FERRUM_TP>1 not supported on the Backend<B> path. \
1520 Run with FERRUM_TP=1 (default) for single-GPU inference.",
1521 ));
1522 }
1523 let _ = ferrum_models::loader::QuantizeConfig::from_model_dir(&model_dir_path);
1529
1530 let (qcfg, moe_cfg): (
1535 ferrum_models::models::LlamaFamilyConfig,
1536 Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1537 ) = match arch {
1538 ferrum_models::Architecture::Qwen3Moe => {
1539 info!("Loading Qwen3-MoE via Qwen3MoeModel (safetensors GPTQ)");
1540 let mc = ferrum_models::moe_config::Qwen3MoeConfig::from_def(&model_def)?;
1541 (mc.base.clone(), Some(mc))
1545 }
1546 ferrum_models::Architecture::Qwen3 => {
1547 info!("Loading Qwen3 via LlamaFamilyModel (QK-norm on)");
1548 (
1549 ferrum_models::models::LlamaFamilyConfig::qwen3_from_def(&model_def),
1550 None,
1551 )
1552 }
1553 ferrum_models::Architecture::Qwen2 => {
1554 info!("Loading Qwen2 via LlamaFamilyModel");
1555 (
1556 ferrum_models::models::LlamaFamilyConfig::qwen2_from_def(&model_def),
1557 None,
1558 )
1559 }
1560 ferrum_models::Architecture::Mistral => {
1561 info!("Loading Mistral via LlamaFamilyModel (sliding_window from config)");
1562 (
1563 ferrum_models::models::LlamaFamilyConfig::mistral_from_def(&model_def),
1564 None,
1565 )
1566 }
1567 ferrum_models::Architecture::Gemma3 => {
1568 info!(
1569 "Loading Gemma3 via LlamaFamilyModel (5:1 SWA, dual rope, GeGLU, \
1570 sandwich norms)"
1571 );
1572 (
1573 ferrum_models::models::LlamaFamilyConfig::gemma3_from_def(&model_def),
1574 None,
1575 )
1576 }
1577 _ => {
1578 info!("Loading Llama via LlamaFamilyModel");
1579 (
1580 ferrum_models::models::LlamaFamilyConfig::llama_from_def(&model_def),
1581 None,
1582 )
1583 }
1584 };
1585
1586 let model_info =
1587 model_def.to_model_info(config.engine_config.model.model_id.to_string());
1588 let llama_layer_split_plan =
1589 resolve_llama_layer_split_plan(config, qcfg.num_layers)?;
1590 let llama_layer_split_pipeline_mode = llama_layer_split_plan
1591 .as_ref()
1592 .map(|plan| resolve_llama_layer_split_pipeline_mode(config, plan.stages.len()))
1593 .transpose()?;
1594
1595 use ferrum_interfaces::kv_dtype::KvFp16;
1601 #[cfg(feature = "cuda")]
1602 use ferrum_interfaces::kv_dtype::KvInt8;
1603 use ferrum_types::KvCacheDtype;
1604 let kv_dtype = config.engine_config.kv_cache.dtype;
1605 let llm: Box<dyn ferrum_models::common::DecoderOnlyLLM> =
1606 match (&config.device, kv_dtype) {
1607 (Device::CPU, KvCacheDtype::Fp16) => {
1608 info!(" Backend: CPU, KV: fp16");
1609 build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
1610 arch,
1611 qcfg,
1612 moe_cfg,
1613 &model_path,
1614 llama_layer_split_plan,
1615 llama_layer_split_pipeline_mode,
1616 )?
1617 }
1618 #[cfg(any(target_os = "macos", target_os = "ios"))]
1619 (Device::Metal, KvCacheDtype::Fp16) => {
1620 #[cfg(feature = "metal")]
1621 {
1622 let dtype_hint = RegistryRuntimeEnv::from_runtime_knobs(
1626 &config.engine_config.runtime,
1627 )
1628 .metal_dtype_hint();
1629 info!(" Backend: Metal (weights {}), KV: fp16", dtype_hint);
1630 build_llm::<ferrum_kernels::backend::metal::MetalBackend, KvFp16>(
1631 arch,
1632 qcfg,
1633 moe_cfg,
1634 &model_path,
1635 llama_layer_split_plan,
1636 llama_layer_split_pipeline_mode,
1637 )?
1638 }
1639 #[cfg(not(feature = "metal"))]
1640 {
1641 return Err(FerrumError::device(
1642 "Metal requested but 'metal' feature not enabled",
1643 ));
1644 }
1645 }
1646 (Device::CUDA(_), KvCacheDtype::Fp16) => {
1647 #[cfg(feature = "cuda")]
1648 {
1649 info!(" Backend: CUDA, KV: fp16");
1650 build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvFp16>(
1651 arch,
1652 qcfg,
1653 moe_cfg,
1654 &model_path,
1655 llama_layer_split_plan,
1656 llama_layer_split_pipeline_mode,
1657 )?
1658 }
1659 #[cfg(not(feature = "cuda"))]
1660 {
1661 return Err(FerrumError::device(
1662 "CUDA requested but 'cuda' feature not enabled",
1663 ));
1664 }
1665 }
1666 (Device::CUDA(_), KvCacheDtype::Int8) => {
1667 #[cfg(feature = "cuda")]
1668 {
1669 if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1676 return Err(FerrumError::unsupported(
1677 "INT8 KV cache is not yet wired through Qwen3MoeModel \
1678 (LlamaFamilyModel-only in PR C). Use --kv-dtype fp16 \
1679 for MoE models or wait for the follow-up PR.",
1680 ));
1681 }
1682 info!(" Backend: CUDA, KV: int8 (paged, vLLM-style)");
1683 build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvInt8>(
1684 arch,
1685 qcfg,
1686 moe_cfg,
1687 &model_path,
1688 llama_layer_split_plan,
1689 llama_layer_split_pipeline_mode,
1690 )?
1691 }
1692 #[cfg(not(feature = "cuda"))]
1693 {
1694 return Err(FerrumError::device(
1695 "CUDA requested but 'cuda' feature not enabled",
1696 ));
1697 }
1698 }
1699 (dev, dt) => {
1700 return Err(FerrumError::unsupported(format!(
1701 "(device={dev:?}, kv_dtype={dt:?}) not implemented — \
1702 see docs/dim5-model-wireup-plan.md"
1703 )));
1704 }
1705 };
1706
1707 Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)))
1708 }
1709 ferrum_models::Architecture::Bert => {
1710 info!("Using BERT executor for embeddings");
1711 let executor = ferrum_models::BertModelExecutor::from_path(
1712 &model_path,
1713 &model_def,
1714 legacy_candle_device()?,
1715 )
1716 .await?;
1717
1718 Ok(Arc::new(executor))
1719 }
1720 ferrum_models::Architecture::Clip => {
1721 info!("Using CLIP executor for multimodal embeddings");
1722 let executor = ferrum_models::ClipModelExecutor::from_path(
1723 &model_path,
1724 legacy_candle_device()?,
1725 dtype,
1726 )?;
1727 Ok(Arc::new(executor))
1728 }
1729 ferrum_models::Architecture::Whisper => {
1730 info!("Using Whisper executor for ASR");
1731 let executor = ferrum_models::WhisperModelExecutor::from_path(
1732 &model_path,
1733 legacy_candle_device()?,
1734 dtype,
1735 )?;
1736 Ok(Arc::new(executor))
1737 }
1738 _ => Err(FerrumError::model(format!(
1739 "Architecture {:?} not supported",
1740 model_def.architecture
1741 ))),
1742 }
1743 }
1744
1745 fn metadata(&self) -> ComponentMetadata {
1746 ComponentMetadata {
1747 name: "llm".to_string(),
1748 version: "0.2.0".to_string(),
1749 description: "LLM executor (LlamaFamily / Qwen3MoE via Backend<B>; \
1750 BERT / CLIP / Whisper via candle)"
1751 .to_string(),
1752 supported_devices: cpu_cuda_and_optional_metal_devices(),
1753 capabilities: vec![
1754 "llama".to_string(),
1755 "qwen2".to_string(),
1756 "qwen3".to_string(),
1757 "qwen3_moe".to_string(),
1758 "mistral".to_string(),
1759 "bert".to_string(),
1760 "clip".to_string(),
1761 "whisper".to_string(),
1762 "safetensors".to_string(),
1763 "gguf".to_string(),
1764 "fp16".to_string(),
1765 "fp32".to_string(),
1766 ],
1767 }
1768 }
1769}
1770
1771pub struct StubTokenizer {
1777 vocab_size: usize,
1778 info: ferrum_interfaces::TokenizerInfo,
1779}
1780
1781impl StubTokenizer {
1782 pub fn new(vocab_size: usize) -> Self {
1784 let info = ferrum_interfaces::TokenizerInfo {
1785 tokenizer_type: ferrum_interfaces::tokenizer::TokenizerType::BPE,
1786 vocab_size,
1787 special_tokens: ferrum_types::SpecialTokens::default(),
1788 supports_incremental: false,
1789 supports_chat_template: false,
1790 max_token_length: None,
1791 model_name: Some("stub".into()),
1792 };
1793
1794 Self { vocab_size, info }
1795 }
1796}
1797
1798impl std::fmt::Debug for StubTokenizer {
1799 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1800 f.debug_struct("StubTokenizer")
1801 .field("vocab_size", &self.vocab_size)
1802 .finish()
1803 }
1804}
1805
1806impl Tokenizer for StubTokenizer {
1807 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<ferrum_types::TokenId>> {
1808 let tokens: Vec<ferrum_types::TokenId> = text
1809 .split_whitespace()
1810 .enumerate()
1811 .map(|(i, _)| ferrum_types::TokenId::new((i % self.vocab_size) as u32))
1812 .collect();
1813
1814 Ok(if tokens.is_empty() {
1815 vec![ferrum_types::TokenId::new(0)]
1816 } else {
1817 tokens
1818 })
1819 }
1820
1821 fn decode(&self, tokens: &[ferrum_types::TokenId], _skip_special: bool) -> Result<String> {
1822 Ok(tokens
1823 .iter()
1824 .map(|t| format!("token_{}", t.get()))
1825 .collect::<Vec<_>>()
1826 .join(" "))
1827 }
1828
1829 fn decode_incremental(
1830 &self,
1831 _prev: &[ferrum_types::TokenId],
1832 next: ferrum_types::TokenId,
1833 ) -> Result<String> {
1834 Ok(format!("token_{} ", next.get()))
1835 }
1836
1837 fn vocab_size(&self) -> usize {
1838 self.vocab_size
1839 }
1840
1841 fn special_tokens(&self) -> &ferrum_types::SpecialTokens {
1842 &self.info.special_tokens
1843 }
1844
1845 fn token_id(&self, _text: &str) -> Option<ferrum_types::TokenId> {
1846 Some(ferrum_types::TokenId::new(0))
1847 }
1848
1849 fn token_text(&self, _token_id: ferrum_types::TokenId) -> Option<&str> {
1850 None
1851 }
1852
1853 fn info(&self) -> ferrum_interfaces::TokenizerInfo {
1854 self.info.clone()
1855 }
1856}
1857
1858static GLOBAL_REGISTRY: OnceLock<Arc<ComponentRegistry>> = OnceLock::new();
1864
1865pub fn global_registry() -> Arc<ComponentRegistry> {
1867 GLOBAL_REGISTRY
1868 .get_or_init(|| {
1869 info!("Initializing global component registry");
1870 Arc::new(ComponentRegistry::with_defaults())
1871 })
1872 .clone()
1873}
1874
1875pub fn set_global_registry(registry: Arc<ComponentRegistry>) -> Result<()> {
1877 GLOBAL_REGISTRY
1878 .set(registry)
1879 .map_err(|_| FerrumError::internal("Global registry already initialized"))
1880}
1881
1882#[cfg(test)]
1887mod tests {
1888 use super::*;
1889 use safetensors::tensor::{serialize_to_file, Dtype, TensorView};
1890 use std::path::{Path, PathBuf};
1891
1892 fn unique_test_dir(name: &str) -> PathBuf {
1893 let mut dir = std::env::temp_dir();
1894 dir.push(format!(
1895 "ferrum-{name}-{}-{}",
1896 std::process::id(),
1897 std::time::SystemTime::now()
1898 .duration_since(std::time::UNIX_EPOCH)
1899 .unwrap()
1900 .as_nanos()
1901 ));
1902 std::fs::create_dir_all(&dir).unwrap();
1903 dir
1904 }
1905
1906 fn write_qwen35_fixture_config(dir: &Path) {
1907 let config = serde_json::json!({
1908 "architectures": ["Qwen3_5ForConditionalGeneration"],
1909 "model_type": "qwen3_5",
1910 "vocab_size": 3,
1911 "max_position_embeddings": 16,
1912 "rms_norm_eps": 1e-6,
1913 "rope_theta": 10000.0,
1914 "tie_word_embeddings": false,
1915 "text_config": {
1916 "model_type": "qwen3_5_text",
1917 "hidden_size": 2,
1918 "intermediate_size": 2,
1919 "num_hidden_layers": 2,
1920 "layer_types": ["linear_attention", "full_attention"],
1921 "linear_num_key_heads": 1,
1922 "linear_num_value_heads": 1,
1923 "linear_key_head_dim": 1,
1924 "linear_value_head_dim": 1,
1925 "linear_conv_kernel_dim": 1,
1926 "mamba_ssm_dtype": "float32",
1927 "head_dim": 2,
1928 "num_attention_heads": 1,
1929 "num_key_value_heads": 1,
1930 "vocab_size": 3,
1931 "max_position_embeddings": 16,
1932 "tie_word_embeddings": false
1933 }
1934 });
1935 std::fs::write(
1936 dir.join("config.json"),
1937 serde_json::to_string_pretty(&config).unwrap(),
1938 )
1939 .unwrap();
1940 }
1941
1942 fn write_qwen35_fixture_safetensors(dir: &Path) {
1943 let tensors: Vec<(String, Vec<f32>)> = vec![
1944 (
1945 "model.embed_tokens.weight".to_string(),
1946 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1947 ),
1948 ("model.norm.weight".to_string(), vec![0.0, 0.0]),
1949 (
1950 "model.lm_head.weight".to_string(),
1951 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1952 ),
1953 (
1954 "model.layers.0.input_layernorm.weight".to_string(),
1955 vec![0.0, 0.0],
1956 ),
1957 (
1958 "model.layers.0.post_attention_layernorm.weight".to_string(),
1959 vec![0.0, 0.0],
1960 ),
1961 (
1962 "model.layers.0.linear_attn.in_proj_qkv.weight".to_string(),
1963 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
1964 ),
1965 (
1966 "model.layers.0.linear_attn.in_proj_z.weight".to_string(),
1967 vec![1.0, -1.0],
1968 ),
1969 (
1970 "model.layers.0.linear_attn.in_proj_b.weight".to_string(),
1971 vec![0.5, 0.25],
1972 ),
1973 (
1974 "model.layers.0.linear_attn.in_proj_a.weight".to_string(),
1975 vec![-0.25, 0.75],
1976 ),
1977 (
1978 "model.layers.0.linear_attn.conv1d.weight".to_string(),
1979 vec![1.0, 1.0, 1.0],
1980 ),
1981 ("model.layers.0.linear_attn.A_log".to_string(), vec![0.0]),
1982 ("model.layers.0.linear_attn.dt_bias".to_string(), vec![0.0]),
1983 (
1984 "model.layers.0.linear_attn.norm.weight".to_string(),
1985 vec![1.0],
1986 ),
1987 (
1988 "model.layers.0.linear_attn.out_proj.weight".to_string(),
1989 vec![1.0, -0.5],
1990 ),
1991 (
1992 "model.layers.0.mlp.gate_proj.weight".to_string(),
1993 vec![0.2, 0.1, -0.1, 0.3],
1994 ),
1995 (
1996 "model.layers.0.mlp.up_proj.weight".to_string(),
1997 vec![0.4, -0.2, 0.3, 0.5],
1998 ),
1999 (
2000 "model.layers.0.mlp.down_proj.weight".to_string(),
2001 vec![1.0, 0.0, 0.0, 1.0],
2002 ),
2003 (
2004 "model.layers.1.input_layernorm.weight".to_string(),
2005 vec![0.0, 0.0],
2006 ),
2007 (
2008 "model.layers.1.post_attention_layernorm.weight".to_string(),
2009 vec![0.0, 0.0],
2010 ),
2011 (
2012 "model.layers.1.self_attn.q_proj.weight".to_string(),
2013 vec![1.0, 0.0, 0.0, 1.0],
2014 ),
2015 (
2016 "model.layers.1.self_attn.k_proj.weight".to_string(),
2017 vec![0.5, 0.0, 0.0, 0.5],
2018 ),
2019 (
2020 "model.layers.1.self_attn.v_proj.weight".to_string(),
2021 vec![1.0, 1.0, -0.5, 0.5],
2022 ),
2023 (
2024 "model.layers.1.self_attn.o_proj.weight".to_string(),
2025 vec![1.0, 0.0, 0.0, 1.0],
2026 ),
2027 (
2028 "model.layers.1.self_attn.q_norm.weight".to_string(),
2029 vec![1.0, 1.0],
2030 ),
2031 (
2032 "model.layers.1.self_attn.k_norm.weight".to_string(),
2033 vec![1.0, 1.0],
2034 ),
2035 (
2036 "model.layers.1.mlp.gate_proj.weight".to_string(),
2037 vec![-0.2, 0.2, 0.1, 0.3],
2038 ),
2039 (
2040 "model.layers.1.mlp.up_proj.weight".to_string(),
2041 vec![0.25, 0.5, -0.3, 0.4],
2042 ),
2043 (
2044 "model.layers.1.mlp.down_proj.weight".to_string(),
2045 vec![0.5, 0.25, -0.2, 0.75],
2046 ),
2047 ];
2048 let views = tensors
2049 .into_iter()
2050 .map(|(name, values)| {
2051 let bytes = values
2052 .iter()
2053 .flat_map(|value| value.to_le_bytes())
2054 .collect::<Vec<_>>()
2055 .into_boxed_slice();
2056 let bytes: &'static [u8] = Box::leak(bytes);
2057 (
2058 name,
2059 TensorView::new(Dtype::F32, vec![values.len()], bytes).unwrap(),
2060 )
2061 })
2062 .collect::<Vec<_>>();
2063 serialize_to_file(
2064 views,
2065 &None::<std::collections::HashMap<String, String>>,
2066 &dir.join("model.safetensors"),
2067 )
2068 .unwrap();
2069 }
2070
2071 fn write_qwen35_fixture_model_dir() -> PathBuf {
2072 let dir = unique_test_dir("qwen35-vnext-route");
2073 write_qwen35_fixture_config(&dir);
2074 write_qwen35_fixture_safetensors(&dir);
2075 dir
2076 }
2077
2078 fn qwen35_fixture_component_config(model_dir: &Path) -> ComponentConfig {
2079 let mut engine_config = EngineConfig::default();
2080 engine_config.backend.device = Device::CPU;
2081 engine_config.backend.backend_options.insert(
2082 "model_path".to_string(),
2083 serde_json::Value::String(model_dir.to_string_lossy().to_string()),
2084 );
2085 ComponentConfig::from_engine_config(&engine_config)
2086 }
2087
2088 fn enable_checkpoint_capture(config: &mut ComponentConfig, output_dir: PathBuf) {
2089 config.engine_config.runtime.vnext_checkpoint_capture =
2090 Some(ferrum_types::VNextCheckpointCaptureConfig {
2091 output_dir,
2092 value_ids: vec!["value.output.logits".to_owned()],
2093 maximum_prefill_waves: 1,
2094 maximum_decode_waves: 0,
2095 capture_product_output: false,
2096 teacher_forcing: None,
2097 });
2098 }
2099
2100 #[test]
2101 fn checkpoint_capture_rejects_stub_fallback() {
2102 let mut config = ComponentConfig::from_engine_config(&EngineConfig::default());
2103 enable_checkpoint_capture(&mut config, PathBuf::from("capture"));
2104
2105 let error = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2106 Ok(_) => panic!("checkpoint capture unexpectedly entered the stub executor"),
2107 Err(error) => error.to_string(),
2108 };
2109 assert!(
2110 error.contains("requires a registered model source"),
2111 "{error}"
2112 );
2113 }
2114
2115 #[test]
2116 fn qwen35_registry_routes_registered_package_before_legacy_loading() {
2117 let dir = write_qwen35_fixture_model_dir();
2118 std::fs::remove_file(dir.join("model.safetensors")).unwrap();
2119 let config = qwen35_fixture_component_config(&dir);
2120
2121 let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2122 Ok(_) => panic!("registered Qwen3.5 package unexpectedly accepted CPU execution"),
2123 Err(err) => err.to_string(),
2124 };
2125
2126 assert!(
2127 err.contains(ferrum_models::vnext::qwen35::EXTERNAL_METADATA_ID),
2128 "{err}"
2129 );
2130 assert!(err.contains("but cpu is not registered"), "{err}");
2131 let _ = std::fs::remove_dir_all(dir);
2132 }
2133
2134 #[test]
2135 fn qwen35_typed_gguf_routes_registered_package_before_legacy_gguf_loading() {
2136 let dir = unique_test_dir("qwen35-typed-gguf-route");
2137 write_qwen35_fixture_config(&dir);
2138 std::fs::write(dir.join("tokenizer.json"), br#"{"version":"1.0"}"#).unwrap();
2139 std::fs::write(
2140 dir.join("tokenizer_config.json"),
2141 br#"{"chat_template":"fixture"}"#,
2142 )
2143 .unwrap();
2144 let gguf = dir.join("model.gguf");
2145 std::fs::write(&gguf, b"not-a-real-gguf").unwrap();
2146 let original = ferrum_interfaces::vnext::OriginalModelSource {
2147 kind: ferrum_interfaces::vnext::ModelSourceKind::LocalDirectory,
2148 location: dir.display().to_string(),
2149 requested_revision: None,
2150 };
2151 let sources = Arc::new(
2152 ProductionModelSourceBundle::open(
2153 &dir,
2154 &dir,
2155 ferrum_models::vnext::ProductionWeightArtifact::gguf_file(&gguf),
2156 ferrum_interfaces::vnext::OriginalModelSources {
2157 semantic: original.clone(),
2158 tokenizer: original.clone(),
2159 weights: ferrum_interfaces::vnext::OriginalModelSource {
2160 kind: ferrum_interfaces::vnext::ModelSourceKind::LocalFile,
2161 location: gguf.display().to_string(),
2162 requested_revision: None,
2163 },
2164 },
2165 )
2166 .unwrap(),
2167 );
2168 let mut engine = EngineConfig::default();
2169 engine.backend.device = Device::CPU;
2170 let mut config = ComponentConfig::from_engine_config(&engine);
2171 config.model_sources = Some(sources);
2172
2173 let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2174 Ok(_) => panic!("registered Qwen3.5 GGUF unexpectedly accepted CPU execution"),
2175 Err(err) => err.to_string(),
2176 };
2177
2178 assert!(
2179 err.contains(ferrum_models::vnext::qwen35::EXTERNAL_METADATA_ID),
2180 "{err}"
2181 );
2182 assert!(err.contains("but cpu is not registered"), "{err}");
2183 assert!(!err.contains("GGUF"), "{err}");
2184 let _ = std::fs::remove_dir_all(dir);
2185 }
2186
2187 #[test]
2188 fn test_registry_creation() {
2189 let registry = ComponentRegistry::new();
2190 assert!(registry.list_tokenizers().is_empty());
2191 }
2192
2193 #[test]
2194 fn test_registry_with_defaults() {
2195 let registry = ComponentRegistry::with_defaults();
2196 assert!(registry.list_tokenizers().contains(&"stub".to_string()));
2197 assert!(registry
2198 .list_samplers()
2199 .contains(&"multinomial".to_string()));
2200 assert!(registry.list_schedulers().contains(&"fifo".to_string()));
2201 }
2202
2203 #[test]
2204 fn test_component_config() {
2205 let mut options = HashMap::new();
2206 options.insert("test_key".to_string(), serde_json::json!("test_value"));
2207
2208 let config = ComponentConfig {
2209 engine_config: EngineConfig::default(),
2210 device: Device::CPU,
2211 component_options: options,
2212 model_sources: None,
2213 prepared_model: None,
2214 };
2215
2216 assert_eq!(
2217 config.get_string_option("test_key"),
2218 Some("test_value".to_string())
2219 );
2220 assert!(config.get_string_option("missing").is_none());
2221 }
2222
2223 fn layer_split_component_config_for_device(device: usize) -> ComponentConfig {
2224 let mut config = EngineConfig::default();
2225 config.backend.device = Device::CUDA(device);
2226 config.backend.backend_options.insert(
2227 "selected_distributed_strategy".to_string(),
2228 serde_json::Value::String("layer_split".to_string()),
2229 );
2230 config.backend.backend_options.insert(
2231 "selected_gpu_devices".to_string(),
2232 serde_json::json!([0, 1]),
2233 );
2234 config.backend.backend_options.insert(
2235 "selected_layer_split_plan".to_string(),
2236 serde_json::Value::String(
2237 "stage0:cuda:0:layers=0-39;stage1:cuda:1:layers=40-79".to_string(),
2238 ),
2239 );
2240 config.backend.backend_options.insert(
2241 "selected_layer_split_stages".to_string(),
2242 serde_json::json!([
2243 {"stage": 0, "device": 0, "layer_start": 0, "layer_end": 39},
2244 {"stage": 1, "device": 1, "layer_start": 40, "layer_end": 79}
2245 ]),
2246 );
2247 ComponentConfig::from_engine_config(&config)
2248 }
2249
2250 #[test]
2251 fn resolves_llama_stage_config_for_selected_cuda_device() {
2252 let config = layer_split_component_config_for_device(1);
2253
2254 let stage = resolve_llama_layer_stage_config(&config, 80)
2255 .unwrap()
2256 .unwrap();
2257
2258 assert_eq!(stage.source_layers, 40..80);
2259 assert!(!stage.load_embedding);
2260 assert!(stage.load_lm_head);
2261 }
2262
2263 #[test]
2264 fn resolves_layer_split_pipeline_mode_default_and_explicit_batch() {
2265 let mut config = layer_split_component_config_for_device(0);
2266
2267 assert_eq!(
2268 resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2269 ferrum_models::models::LlamaPipelineMode::Overlapped
2270 );
2271
2272 config.component_options.insert(
2273 "layer_split_pipeline_mode".to_string(),
2274 serde_json::Value::String("batch".to_string()),
2275 );
2276 assert_eq!(
2277 resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2278 ferrum_models::models::LlamaPipelineMode::Batch
2279 );
2280 }
2281
2282 #[test]
2283 fn rejects_llama_stage_config_when_plan_layer_count_mismatches_model() {
2284 let config = layer_split_component_config_for_device(0);
2285
2286 let err = resolve_llama_layer_stage_config(&config, 81)
2287 .unwrap_err()
2288 .to_string();
2289
2290 assert!(err.contains("covers 80 layers but model has 81"));
2291 }
2292
2293 #[test]
2294 fn build_llm_rejects_layer_split_before_weight_open_without_device_scope() {
2295 use ferrum_interfaces::kv_dtype::KvFp16;
2296 use ferrum_models::models::LlamaFamilyConfig;
2297
2298 let qcfg = LlamaFamilyConfig {
2299 hidden_size: 1,
2300 intermediate_size: 1,
2301 num_heads: 1,
2302 num_kv_heads: 1,
2303 head_dim: 1,
2304 num_layers: 2,
2305 vocab_size: 1,
2306 max_seq_len: 8,
2307 rms_norm_eps: 1e-5,
2308 rope_theta: 10_000.0,
2309 rope_scaling: None,
2310 rope_interleaved: false,
2311 has_qk_norm: false,
2312 sliding_window: 0,
2313 ..Default::default()
2314 };
2315 let plan = crate::layer_split::parse_layer_split_plan(
2316 "stage0:cuda:0:layers=0-0;stage1:cuda:1:layers=1-1",
2317 )
2318 .unwrap();
2319
2320 let err = match build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
2321 ferrum_models::Architecture::Llama,
2322 qcfg,
2323 None,
2324 "/missing/model/path",
2325 Some(plan),
2326 Some(ferrum_models::models::LlamaPipelineMode::Overlapped),
2327 ) {
2328 Ok(_) => panic!("layer_split build unexpectedly succeeded"),
2329 Err(err) => err.to_string(),
2330 };
2331
2332 assert!(err.contains("device-scoped execution"));
2333 assert!(err.contains("refusing to silently use the default device"));
2334 }
2335
2336 #[test]
2337 fn test_registry_runtime_env_parses_model_dtype_and_tp() {
2338 let env = RegistryRuntimeEnv::from_env_vars([
2339 ("FERRUM_MODEL_PATH", "/models/qwen"),
2340 ("FERRUM_METAL_DTYPE", "fp16"),
2341 ("FERRUM_DTYPE", "fp32"),
2342 ("FERRUM_TP", "4"),
2343 ]);
2344
2345 assert_eq!(env.model_path.as_deref(), Some("/models/qwen"));
2346 assert_eq!(env.metal_dtype.as_deref(), Some("fp16"));
2347 assert_eq!(env.dtype.as_deref(), Some("fp32"));
2348 assert_eq!(env.tp, 4);
2349 }
2350
2351 #[test]
2352 fn test_registry_runtime_env_defaults_invalid_tp_and_cuda_dtype() {
2353 let env = RegistryRuntimeEnv::from_env_vars([
2354 ("FERRUM_DTYPE", "not-a-dtype"),
2355 ("FERRUM_TP", "not-a-number"),
2356 ]);
2357
2358 assert_eq!(env.model_path(), None);
2359 assert_eq!(env.tp, 0);
2360 assert_eq!(
2361 env.dtype_for_device(&Device::CUDA(0)),
2362 candle_core::DType::F16
2363 );
2364 assert_eq!(env.dtype_for_device(&Device::CPU), candle_core::DType::F32);
2365 }
2366
2367 #[cfg(any(target_os = "macos", target_os = "ios"))]
2368 #[test]
2369 fn test_registry_runtime_env_metal_dtype_precedence() {
2370 let env = RegistryRuntimeEnv::from_env_vars([
2371 ("FERRUM_METAL_DTYPE", "fp16"),
2372 ("FERRUM_DTYPE", "fp32"),
2373 ]);
2374
2375 assert_eq!(
2376 env.dtype_for_device(&Device::Metal),
2377 candle_core::DType::F16
2378 );
2379 assert_eq!(env.metal_dtype_hint(), "fp16");
2380 }
2381
2382 #[test]
2383 fn test_stub_tokenizer() {
2384 let tokenizer = StubTokenizer::new(100);
2385 assert_eq!(tokenizer.vocab_size(), 100);
2386
2387 let tokens = tokenizer.encode("hello world", false).unwrap();
2388 assert!(!tokens.is_empty());
2389
2390 let text = tokenizer.decode(&tokens, false).unwrap();
2391 assert!(text.contains("token_"));
2392 }
2393
2394 #[test]
2395 fn test_greedy_sampler() {
2396 use rand::SeedableRng;
2397
2398 let sampler = GreedySampler;
2399 let mut rng = rand::rngs::StdRng::seed_from_u64(42);
2400
2401 let logits = vec![0.1, 0.5, 0.3, 0.9, 0.2];
2402 let token = sampler.sample(&logits, &mut rng).unwrap();
2403
2404 assert_eq!(token.get(), 3); }
2406}