1use async_trait::async_trait;
14use ferrum_interfaces::{
15 KvCacheManager, ModelExecutor, Sampler, SchedulerInterface as Scheduler, Tokenizer,
16};
17use ferrum_models::vnext::{DefinedProductionModel, 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 defined_model: Option<Arc<DefinedProductionModel>>,
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 defined_model: Option<Arc<DefinedProductionModel>>,
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 defined_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;
959
960fn validate_fixed_storage_kv_dtype(dtype: ferrum_types::KvCacheDtype, loader: &str) -> Result<()> {
964 if dtype != ferrum_types::KvCacheDtype::Fp16 {
965 return Err(FerrumError::unsupported(format!(
966 "{loader} does not support KV dtype {}; use --kv-dtype fp16 for its existing storage behavior",
967 dtype.as_str()
968 )));
969 }
970 Ok(())
971}
972
973fn resolve_llama_layer_split_plan(
974 config: &ComponentConfig,
975 num_layers: usize,
976) -> Result<Option<crate::layer_split::ParsedLayerSplitPlan>> {
977 if config
978 .get_string_option("selected_distributed_strategy")
979 .as_deref()
980 != Some("layer_split")
981 {
982 return Ok(None);
983 }
984
985 let selected = config
986 .get_option::<Vec<usize>>("selected_gpu_devices")
987 .unwrap_or_default();
988 let plan_raw = config.get_string_option("selected_layer_split_plan");
989 let parsed_plan =
990 if let Some(stages) = config.component_options.get("selected_layer_split_stages") {
991 crate::layer_split::parse_layer_split_stage_documents(stages)?
992 } else {
993 let plan_raw = plan_raw.as_deref().ok_or_else(|| {
994 FerrumError::config(
995 "selected_distributed_strategy=layer_split requires selected_layer_split_plan",
996 )
997 })?;
998 crate::layer_split::parse_layer_split_plan(plan_raw)?
999 };
1000 crate::layer_split::validate_layer_split_plan_for_devices(&parsed_plan, &selected)?;
1001 if parsed_plan.total_layers() != num_layers {
1002 return Err(FerrumError::config(format!(
1003 "selected_layer_split_plan covers {} layers but model has {num_layers}",
1004 parsed_plan.total_layers()
1005 )));
1006 }
1007 Ok(Some(parsed_plan))
1008}
1009
1010fn resolve_llama_layer_split_pipeline_mode(
1011 config: &ComponentConfig,
1012 stage_count: usize,
1013) -> Result<ferrum_models::models::LlamaPipelineMode> {
1014 let Some(mode) = config.get_string_option("layer_split_pipeline_mode") else {
1015 return Ok(ferrum_models::models::LlamaPipelineMode::default_for_stage_count(stage_count));
1016 };
1017 let mode = ferrum_models::models::LlamaPipelineMode::from_config_value(&mode)?;
1018 if mode == ferrum_models::models::LlamaPipelineMode::Overlapped && stage_count != 2 {
1019 return Err(FerrumError::config(
1020 "layer_split_pipeline_mode=overlapped requires exactly two pipeline stages",
1021 ));
1022 }
1023 Ok(mode)
1024}
1025
1026#[cfg(test)]
1027fn resolve_llama_layer_stage_config(
1028 config: &ComponentConfig,
1029 num_layers: usize,
1030) -> Result<Option<ferrum_models::models::llama_family::LlamaFamilyLayerStageConfig>> {
1031 let Some(parsed_plan) = resolve_llama_layer_split_plan(config, num_layers)? else {
1032 return Ok(None);
1033 };
1034 let device_id = match &config.device {
1035 Device::CUDA(device_id) => *device_id,
1036 other => {
1037 return Err(FerrumError::unsupported(format!(
1038 "selected_distributed_strategy=layer_split requires a CUDA stage device, got {other:?}",
1039 )));
1040 }
1041 };
1042 parsed_plan
1043 .llama_stage_config_for_device(device_id)
1044 .map(Some)
1045}
1046
1047fn build_llm<B, K>(
1057 arch: ferrum_models::Architecture,
1058 qcfg: ferrum_models::models::LlamaFamilyConfig,
1059 moe_cfg: Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1060 model_path: &str,
1061 llama_layer_split_plan: Option<crate::layer_split::ParsedLayerSplitPlan>,
1062 llama_layer_split_pipeline_mode: Option<ferrum_models::models::LlamaPipelineMode>,
1063) -> Result<Box<dyn ferrum_models::common::DecoderOnlyLLM>>
1064where
1065 B: ferrum_kernels::backend::MoeLlmBackend,
1066 K: ferrum_kernels::backend::KvLayer<B>,
1067 ferrum_models::models::LlamaFamilyModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1068 ferrum_models::models::LlamaFamilyPipelineModel<B, K>: ferrum_models::common::DecoderOnlyLLM,
1069{
1070 if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1071 if llama_layer_split_plan.is_some() {
1072 return Err(FerrumError::unsupported(
1073 "CUDA layer_split stage loading is wired only for Llama-family dense models; \
1074 Qwen3MoeModel requires a separate MoE stage loader.",
1075 ));
1076 }
1077 let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1078 let mc = moe_cfg.ok_or_else(|| {
1079 FerrumError::internal(
1080 "Qwen3Moe arch reached build_llm without Qwen3MoeConfig (caller bug)",
1081 )
1082 })?;
1083 Ok(Box::new(
1084 ferrum_models::models::Qwen3MoeModel::<B, K>::new_safetensors(mc, &weight_loader)?,
1085 ))
1086 } else {
1087 if llama_layer_split_plan.is_some() && !B::supports_device_ordinal_scope() {
1088 return Err(FerrumError::unsupported(
1089 "selected_distributed_strategy=layer_split requires a backend with \
1090 device-scoped execution; refusing to silently use the default device",
1091 ));
1092 }
1093 let weight_loader = ferrum_quantization::NativeSafetensorsLoader::<B>::open(model_path)?;
1094 if let Some(plan) = llama_layer_split_plan {
1095 let stage_configs = plan.to_llama_stage_configs();
1096 let stage_device_ordinals = plan
1097 .stages
1098 .iter()
1099 .map(|stage| Some(stage.device))
1100 .collect::<Vec<_>>();
1101 let mut stages = Vec::with_capacity(stage_configs.len());
1102 for (idx, (stage_config, device_ordinal)) in stage_configs
1103 .into_iter()
1104 .zip(stage_device_ordinals.iter().copied())
1105 .enumerate()
1106 {
1107 tracing::info!(
1108 "Loading Llama layer_split stage {idx} on backend device {:?}",
1109 device_ordinal
1110 );
1111 let stage = B::with_device_ordinal(device_ordinal, || {
1112 ferrum_models::models::LlamaFamilyModel::<B, K>::new_layer_stage(
1113 qcfg.clone(),
1114 &weight_loader,
1115 stage_config,
1116 )
1117 })?;
1118 stages.push(stage);
1119 }
1120 Ok(Box::new(ferrum_models::models::LlamaFamilyPipelineModel::<
1121 B,
1122 K,
1123 >::new_with_placement(
1124 stages,
1125 ferrum_models::models::LlamaPipelinePlacement::from_backend_device_ordinals(
1126 stage_device_ordinals,
1127 )
1128 .with_pipeline_mode(
1129 llama_layer_split_pipeline_mode.expect("layer split pipeline mode resolved"),
1130 ),
1131 )?))
1132 } else {
1133 Ok(Box::new(
1134 ferrum_models::models::LlamaFamilyModel::<B, K>::new(qcfg, &weight_loader)?,
1135 ))
1136 }
1137 }
1138}
1139
1140fn validate_registered_vnext_backend(
1141 kind: ferrum_models::vnext::ProductionExecutionKind,
1142 device: &Device,
1143 external_metadata_id: &ferrum_interfaces::vnext::ExternalModelMetadataId,
1144) -> Result<()> {
1145 use ferrum_models::vnext::ProductionExecutionKind;
1146
1147 match (kind, device) {
1148 (ProductionExecutionKind::CausalLanguage, Device::CUDA(_)) => {
1149 #[cfg(feature = "cuda")]
1150 return Ok(());
1151 #[cfg(not(feature = "cuda"))]
1152 Err(FerrumError::device(
1153 "registered vNext CUDA composition requires the 'cuda' feature",
1154 ))
1155 }
1156 #[cfg(any(target_os = "macos", target_os = "ios"))]
1157 (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1158 #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1159 return Ok(());
1160 #[cfg(not(all(feature = "metal", any(target_os = "macos", target_os = "ios"))))]
1161 Err(FerrumError::device(
1162 "registered vNext Metal composition requires the 'metal' feature on an Apple platform",
1163 ))
1164 }
1165 (ProductionExecutionKind::CausalLanguage, Device::CPU) => Ok(()),
1166 (kind, device) => Err(FerrumError::unsupported(format!(
1167 "registered vNext model metadata {external_metadata_id} requires a {kind:?} backend composition, but {device} is not registered"
1168 ))),
1169 }
1170}
1171
1172fn create_registered_vnext_executor(
1173 config: &ComponentConfig,
1174 model_path: &std::path::Path,
1175 sources: Option<Arc<ProductionModelSourceBundle>>,
1176 defined_model: Option<Arc<DefinedProductionModel>>,
1177 registration: ferrum_models::vnext::RegisteredProductionModel,
1178) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1179 use ferrum_models::vnext::ProductionExecutionKind;
1180
1181 validate_registered_vnext_backend(
1182 registration.execution_kind(),
1183 &config.device,
1184 registration.external_metadata_id(),
1185 )?;
1186
1187 let _defined_model_reused = defined_model.is_some();
1188 if let (Some(prepared), Some(sources)) = (defined_model.as_ref(), sources.as_ref()) {
1189 if !Arc::ptr_eq(prepared.sources(), sources) {
1190 return Err(FerrumError::model(
1191 "prepared product model and component sources do not share one source lease",
1192 ));
1193 }
1194 }
1195 let prepared = match defined_model {
1196 Some(prepared) => {
1197 if prepared.definition().external_metadata_id() != registration.external_metadata_id() {
1198 return Err(FerrumError::model(format!(
1199 "prepared product metadata {} differs from registered metadata {}",
1200 prepared.definition().external_metadata_id(),
1201 registration.external_metadata_id()
1202 )));
1203 }
1204 prepared
1205 }
1206 None => Arc::new(match sources {
1207 Some(sources) => registration.define_from_sources(sources)?,
1208 None => registration.define(model_path)?,
1209 }),
1210 };
1211
1212 match (registration.execution_kind(), &config.device) {
1213 (ProductionExecutionKind::CausalLanguage, Device::CUDA(ordinal)) => {
1214 #[cfg(feature = "cuda")]
1215 {
1216 let family = prepared.definition();
1217 let family_fingerprint = family
1218 .fingerprint()
1219 .map_err(|error| FerrumError::model(error.to_string()))?;
1220 info!(
1221 external_metadata_id = %registration.external_metadata_id(),
1222 family_id = %family.family_id(),
1223 family_fingerprint,
1224 defined_model_reused = _defined_model_reused,
1225 backend = "cuda",
1226 "Resolving a defined model against the actual vNext runtime"
1227 );
1228 let device_id = ferrum_interfaces::vnext::DeviceId::new(format!(
1229 "device.cuda.{ordinal}"
1230 ))
1231 .map_err(|error| FerrumError::device(error.to_string()))?;
1232 let composition =
1233 ferrum_kernels::backend::cuda::vnext_ops::CudaVNextComposition::create(
1234 *ordinal,
1235 device_id,
1236 crate::product_composition::cuda_attention_policy_for_kv(
1237 config.engine_config.runtime.attention_execution_policy,
1238 config.engine_config.kv_cache.dtype,
1239 )?,
1240 )
1241 .map_err(|error| {
1242 FerrumError::device(format!("create vNext CUDA runtime: {error}"))
1243 })?;
1244 let (
1245 runtime,
1246 operation_registry,
1247 weight_materializers,
1248 catalog,
1249 ) = composition.into_parts();
1250 let executor = crate::product_composition::create_vnext_executor(
1251 &config.engine_config,
1252 prepared.as_ref(),
1253 runtime,
1254 operation_registry,
1255 weight_materializers,
1256 catalog,
1257 ferrum_kernels::backend::cuda::vnext_ops::cuda_weight_materializer_selection,
1258 )?;
1259 info!(
1260 resolved_plan_fingerprint = executor
1261 .resolved_model_plan()
1262 .map(|plan| plan.fingerprint())
1263 .unwrap_or("missing"),
1264 "Resolved product model plan is authoritative for vNext execution"
1265 );
1266 Ok(Arc::new(executor))
1267 }
1268 #[cfg(not(feature = "cuda"))]
1269 {
1270 let _ = (ordinal, model_path, prepared);
1271 Err(FerrumError::device(
1272 "registered vNext CUDA composition requires the 'cuda' feature",
1273 ))
1274 }
1275 }
1276 #[cfg(all(feature = "metal", any(target_os = "macos", target_os = "ios")))]
1277 (ProductionExecutionKind::CausalLanguage, Device::Metal) => {
1278 let family = prepared.definition();
1279 let family_fingerprint = family
1280 .fingerprint()
1281 .map_err(|error| FerrumError::model(error.to_string()))?;
1282 info!(
1283 external_metadata_id = %registration.external_metadata_id(),
1284 family_id = %family.family_id(),
1285 family_fingerprint,
1286 defined_model_reused = _defined_model_reused,
1287 backend = "metal",
1288 "Resolving a defined model against the actual vNext runtime"
1289 );
1290 let device_id = ferrum_interfaces::vnext::DeviceId::new("device.metal.0")
1291 .map_err(|error| FerrumError::device(error.to_string()))?;
1292 let composition =
1293 ferrum_kernels::backend::metal::vnext_ops::MetalVNextComposition::create(
1294 device_id,
1295 )
1296 .map_err(|error| {
1297 FerrumError::device(format!("create vNext Metal runtime: {error}"))
1298 })?;
1299 let (
1300 runtime,
1301 operation_registry,
1302 weight_materializers,
1303 weight_materializer_id,
1304 catalog,
1305 ) = composition.into_parts();
1306 let weight_materializer_selection =
1307 ferrum_interfaces::vnext::WeightMaterializerSelection::exact(
1308 weight_materializer_id,
1309 );
1310 let executor = crate::product_composition::create_vnext_executor(
1311 &config.engine_config,
1312 prepared.as_ref(),
1313 runtime,
1314 operation_registry,
1315 weight_materializers,
1316 catalog,
1317 |_| Ok(weight_materializer_selection.clone()),
1318 )?;
1319 info!(
1320 resolved_plan_fingerprint = executor
1321 .resolved_model_plan()
1322 .map(|plan| plan.fingerprint())
1323 .unwrap_or("missing"),
1324 "Resolved product model plan is authoritative for vNext execution"
1325 );
1326 Ok(Arc::new(executor))
1327 }
1328 (ProductionExecutionKind::CausalLanguage, Device::CPU) => {
1329 let device_id = ferrum_interfaces::vnext::DeviceId::new("device.cpu.0")
1330 .map_err(|error| FerrumError::device(error.to_string()))?;
1331 let composition = ferrum_kernels::backend::cpu::vnext_ops::CpuVNextComposition::for_host(device_id)
1332 .map_err(|error| FerrumError::device(format!("create vNext CPU runtime: {error}")))?;
1333 let (runtime, operation_registry, weight_materializers, weight_materializer_id, catalog) = composition.into_parts()
1334 .map_err(|error| FerrumError::config(error.to_string()))?;
1335 let selection = ferrum_interfaces::vnext::WeightMaterializerSelection::exact(weight_materializer_id);
1336 info!(
1337 external_metadata_id = %registration.external_metadata_id(),
1338 family_id = %prepared.definition().family_id(),
1339 defined_model_reused = _defined_model_reused,
1340 backend = "cpu",
1341 "Resolving a defined model against the actual vNext runtime"
1342 );
1343 let executor = crate::product_composition::create_vnext_executor(
1344 &config.engine_config, prepared.as_ref(), runtime, operation_registry,
1345 weight_materializers, catalog, |_| Ok(selection.clone()),
1346 )?;
1347 info!(
1348 resolved_plan_fingerprint = executor.resolved_model_plan().map(|plan| plan.fingerprint()).unwrap_or("missing"),
1349 "Resolved product model plan is authoritative for vNext execution"
1350 );
1351 Ok(Arc::new(executor))
1352 }
1353 (kind, device) => Err(FerrumError::unsupported(format!(
1354 "registered vNext model metadata {} requires a {kind:?} backend composition, but {device} is not registered",
1355 registration.external_metadata_id()
1356 ))),
1357 }
1358}
1359
1360#[deprecated(note = "use `LlmExecutorFactory` (renamed PR A — Dim 1/3 cleanup)")]
1363pub type CandleExecutorFactory = LlmExecutorFactory;
1364
1365#[async_trait]
1366impl ComponentFactory<Arc<dyn ModelExecutor + Send + Sync>> for LlmExecutorFactory {
1367 async fn create(
1368 &self,
1369 config: &ComponentConfig,
1370 ) -> Result<Arc<dyn ModelExecutor + Send + Sync>> {
1371 use candle_core::{DType, Device as CandleDevice};
1372 use ferrum_models::weight_format::WeightFormat;
1373
1374 let checkpoint_capture_enabled = config
1376 .engine_config
1377 .runtime
1378 .vnext_checkpoint_capture
1379 .is_some();
1380 let model_sources = config.model_sources.clone();
1381 let defined_model = config.defined_model.clone();
1382 let model_path = model_sources
1383 .as_ref()
1384 .map(|sources| sources.weights().path().display().to_string())
1385 .or_else(|| config.get_string_option("model_path"))
1386 .or_else(|| config.engine_config.runtime.model_path.clone());
1387
1388 let model_path = match model_path {
1389 Some(path) => path,
1390 None => {
1391 if checkpoint_capture_enabled {
1392 return Err(FerrumError::unsupported(
1393 "vNext checkpoint capture requires a registered model source",
1394 ));
1395 }
1396 info!("No model path found, falling back to stub executor");
1397 return StubExecutorFactory.create(config).await;
1398 }
1399 };
1400
1401 let weight_fmt = WeightFormat::detect(std::path::Path::new(&model_path))?;
1406 info!(
1407 "Loading model from {} (format: {})",
1408 model_path,
1409 weight_fmt.label()
1410 );
1411
1412 let production_registration = match model_sources.as_ref() {
1421 Some(sources) => Some(ferrum_models::vnext::resolve_registered_model_from_sources(
1422 sources,
1423 )?),
1424 None if !matches!(weight_fmt, WeightFormat::Gguf { .. }) => {
1425 Some(ferrum_models::vnext::resolve_registered_model_from_dir(
1426 std::path::Path::new(&model_path),
1427 )?)
1428 }
1429 None => None,
1430 };
1431 if let Some(production_registration) = production_registration {
1432 match production_registration {
1433 ferrum_models::vnext::ProductionModelRegistration::Registered(registration) => {
1434 return create_registered_vnext_executor(
1435 config,
1436 std::path::Path::new(&model_path),
1437 model_sources,
1438 defined_model,
1439 registration,
1440 );
1441 }
1442 ferrum_models::vnext::ProductionModelRegistration::LegacyRegistered {
1443 external_metadata_id,
1444 } => {
1445 if checkpoint_capture_enabled {
1446 return Err(FerrumError::unsupported(format!(
1447 "vNext checkpoint capture requires a migrated model; metadata {external_metadata_id} is still legacy"
1448 )));
1449 }
1450 info!(
1451 %external_metadata_id,
1452 "Entering the explicitly registered legacy model path"
1453 );
1454 }
1455 }
1456 }
1457 if checkpoint_capture_enabled {
1458 return Err(FerrumError::unsupported(
1459 "vNext checkpoint capture requires a registered vNext model package",
1460 ));
1461 }
1462
1463 if let ferrum_types::NumericalExecutionPolicy::Require(profile) =
1464 &config.engine_config.numerical_execution
1465 {
1466 return Err(FerrumError::unsupported(format!(
1467 "numerical profile {profile} requires a registered vNext model package"
1468 )));
1469 }
1470
1471 if let WeightFormat::Gguf { ref path } = weight_fmt {
1472 validate_fixed_storage_kv_dtype(
1475 config.engine_config.kv_cache.dtype,
1476 "legacy GGUF loader",
1477 )?;
1478 let (llm, model_info) = ferrum_models::gguf_engine_loader::load_gguf_decoder_with_info(
1479 path,
1480 &config.device,
1481 config.engine_config.model.model_id.clone(),
1482 )?;
1483 return Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)));
1484 }
1485
1486 let mut config_manager = ferrum_models::ConfigManager::new();
1490 let model_def = config_manager
1491 .load_from_path(std::path::Path::new(&model_path))
1492 .await?;
1493
1494 info!(
1495 " Architecture: {:?}, Layers: {}, Vocab: {}",
1496 model_def.architecture, model_def.num_hidden_layers, model_def.vocab_size
1497 );
1498
1499 let legacy_candle_device = || -> Result<CandleDevice> {
1506 validate_fixed_storage_kv_dtype(
1507 config.engine_config.kv_cache.dtype,
1508 "legacy Candle executor",
1509 )?;
1510 match &config.device {
1511 Device::CPU => Ok(CandleDevice::Cpu),
1512 #[cfg(feature = "candle-cuda-compat")]
1513 Device::CUDA(id) => CandleDevice::new_cuda(*id)
1514 .map_err(|e| FerrumError::device(format!("CUDA error: {}", e))),
1515 #[cfg(not(feature = "candle-cuda-compat"))]
1516 Device::CUDA(_) => Err(FerrumError::unsupported(
1517 "legacy Candle CUDA executors require the candle-cuda-compat feature",
1518 )),
1519 #[cfg(any(target_os = "macos", target_os = "ios"))]
1520 Device::Metal => CandleDevice::new_metal(0)
1521 .map_err(|e| FerrumError::device(format!("Metal error: {}", e))),
1522 Device::ROCm(_) => Err(FerrumError::device("ROCm not yet supported")),
1523 }
1524 };
1525
1526 let dtype: DType = RegistryRuntimeEnv::from_runtime_knobs(&config.engine_config.runtime)
1535 .dtype_for_device(&config.device);
1536
1537 info!("Building model...");
1539 match model_def.architecture {
1540 arch @ (ferrum_models::Architecture::Llama
1544 | ferrum_models::Architecture::Qwen2
1545 | ferrum_models::Architecture::Qwen3
1546 | ferrum_models::Architecture::Qwen3Moe
1547 | ferrum_models::Architecture::Gemma3
1548 | ferrum_models::Architecture::Mistral) => {
1549 let _loader = ferrum_models::SafeTensorsLoader::new(&model_path);
1550 let model_dir_path: std::path::PathBuf = model_path.clone().into();
1551
1552 let tp_size = config.engine_config.runtime.tp.unwrap_or(0);
1557 if tp_size > 1 {
1558 return Err(FerrumError::unsupported(
1559 "FERRUM_TP>1 not supported on the Backend<B> path. \
1560 Run with FERRUM_TP=1 (default) for single-GPU inference.",
1561 ));
1562 }
1563 let _ = ferrum_models::loader::QuantizeConfig::from_model_dir(&model_dir_path);
1569
1570 let (qcfg, moe_cfg): (
1575 ferrum_models::models::LlamaFamilyConfig,
1576 Option<ferrum_models::moe_config::Qwen3MoeConfig>,
1577 ) = match arch {
1578 ferrum_models::Architecture::Qwen3Moe => {
1579 info!("Loading Qwen3-MoE via Qwen3MoeModel (safetensors GPTQ)");
1580 let mc = ferrum_models::moe_config::Qwen3MoeConfig::from_def(&model_def)?;
1581 (mc.base.clone(), Some(mc))
1585 }
1586 ferrum_models::Architecture::Qwen3 => {
1587 info!("Loading Qwen3 via LlamaFamilyModel (QK-norm on)");
1588 (
1589 ferrum_models::models::LlamaFamilyConfig::qwen3_from_def(&model_def),
1590 None,
1591 )
1592 }
1593 ferrum_models::Architecture::Qwen2 => {
1594 info!("Loading Qwen2 via LlamaFamilyModel");
1595 (
1596 ferrum_models::models::LlamaFamilyConfig::qwen2_from_def(&model_def),
1597 None,
1598 )
1599 }
1600 ferrum_models::Architecture::Mistral => {
1601 info!("Loading Mistral via LlamaFamilyModel (sliding_window from config)");
1602 (
1603 ferrum_models::models::LlamaFamilyConfig::mistral_from_def(&model_def),
1604 None,
1605 )
1606 }
1607 ferrum_models::Architecture::Gemma3 => {
1608 info!(
1609 "Loading Gemma3 via LlamaFamilyModel (5:1 SWA, dual rope, GeGLU, \
1610 sandwich norms)"
1611 );
1612 (
1613 ferrum_models::models::LlamaFamilyConfig::gemma3_from_def(&model_def),
1614 None,
1615 )
1616 }
1617 _ => {
1618 info!("Loading Llama via LlamaFamilyModel");
1619 (
1620 ferrum_models::models::LlamaFamilyConfig::llama_from_def(&model_def),
1621 None,
1622 )
1623 }
1624 };
1625
1626 let model_info =
1627 model_def.to_model_info(config.engine_config.model.model_id.to_string());
1628 let llama_layer_split_plan =
1629 resolve_llama_layer_split_plan(config, qcfg.num_layers)?;
1630 let llama_layer_split_pipeline_mode = llama_layer_split_plan
1631 .as_ref()
1632 .map(|plan| resolve_llama_layer_split_pipeline_mode(config, plan.stages.len()))
1633 .transpose()?;
1634
1635 use ferrum_interfaces::kv_dtype::KvFp16;
1641 #[cfg(feature = "cuda")]
1642 use ferrum_interfaces::kv_dtype::KvInt8;
1643 use ferrum_types::KvCacheDtype;
1644 let kv_dtype = config.engine_config.kv_cache.dtype;
1645 let llm: Box<dyn ferrum_models::common::DecoderOnlyLLM> =
1646 match (&config.device, kv_dtype) {
1647 (Device::CPU, KvCacheDtype::Fp16) => {
1648 info!(" Backend: CPU, KV: fp16");
1649 build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
1650 arch,
1651 qcfg,
1652 moe_cfg,
1653 &model_path,
1654 llama_layer_split_plan,
1655 llama_layer_split_pipeline_mode,
1656 )?
1657 }
1658 #[cfg(any(target_os = "macos", target_os = "ios"))]
1659 (Device::Metal, KvCacheDtype::Fp16) => {
1660 #[cfg(feature = "metal")]
1661 {
1662 let dtype_hint = RegistryRuntimeEnv::from_runtime_knobs(
1666 &config.engine_config.runtime,
1667 )
1668 .metal_dtype_hint();
1669 info!(" Backend: Metal (weights {}), KV: fp16", dtype_hint);
1670 build_llm::<ferrum_kernels::backend::metal::MetalBackend, KvFp16>(
1671 arch,
1672 qcfg,
1673 moe_cfg,
1674 &model_path,
1675 llama_layer_split_plan,
1676 llama_layer_split_pipeline_mode,
1677 )?
1678 }
1679 #[cfg(not(feature = "metal"))]
1680 {
1681 return Err(FerrumError::device(
1682 "Metal requested but 'metal' feature not enabled",
1683 ));
1684 }
1685 }
1686 (Device::CUDA(_), KvCacheDtype::Fp16) => {
1687 #[cfg(feature = "cuda")]
1688 {
1689 info!(" Backend: CUDA, KV: fp16");
1690 build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvFp16>(
1691 arch,
1692 qcfg,
1693 moe_cfg,
1694 &model_path,
1695 llama_layer_split_plan,
1696 llama_layer_split_pipeline_mode,
1697 )?
1698 }
1699 #[cfg(not(feature = "cuda"))]
1700 {
1701 return Err(FerrumError::device(
1702 "CUDA requested but 'cuda' feature not enabled",
1703 ));
1704 }
1705 }
1706 (Device::CUDA(_), KvCacheDtype::Int8) => {
1707 #[cfg(feature = "cuda")]
1708 {
1709 if matches!(arch, ferrum_models::Architecture::Qwen3Moe) {
1716 return Err(FerrumError::unsupported(
1717 "INT8 KV cache is not yet wired through Qwen3MoeModel \
1718 (LlamaFamilyModel-only in PR C). Use --kv-dtype fp16 \
1719 for MoE models or wait for the follow-up PR.",
1720 ));
1721 }
1722 info!(" Backend: CUDA, KV: int8 (paged, vLLM-style)");
1723 build_llm::<ferrum_kernels::backend::cuda::CudaBackend, KvInt8>(
1724 arch,
1725 qcfg,
1726 moe_cfg,
1727 &model_path,
1728 llama_layer_split_plan,
1729 llama_layer_split_pipeline_mode,
1730 )?
1731 }
1732 #[cfg(not(feature = "cuda"))]
1733 {
1734 return Err(FerrumError::device(
1735 "CUDA requested but 'cuda' feature not enabled",
1736 ));
1737 }
1738 }
1739 (dev, dt) => {
1740 return Err(FerrumError::unsupported(format!(
1741 "(device={dev:?}, kv_dtype={dt:?}) not implemented"
1742 )));
1743 }
1744 };
1745
1746 Ok(Arc::new(ferrum_models::LlmExecutor::new(llm, model_info)))
1747 }
1748 ferrum_models::Architecture::Bert => {
1749 info!("Using BERT executor for embeddings");
1750 let executor = ferrum_models::BertModelExecutor::from_path(
1751 &model_path,
1752 &model_def,
1753 legacy_candle_device()?,
1754 )
1755 .await?;
1756
1757 Ok(Arc::new(executor))
1758 }
1759 ferrum_models::Architecture::Clip => {
1760 info!("Using CLIP executor for multimodal embeddings");
1761 let executor = ferrum_models::ClipModelExecutor::from_path(
1762 &model_path,
1763 legacy_candle_device()?,
1764 dtype,
1765 )?;
1766 Ok(Arc::new(executor))
1767 }
1768 ferrum_models::Architecture::Whisper => {
1769 info!("Using Whisper executor for ASR");
1770 let executor = ferrum_models::WhisperModelExecutor::from_path(
1771 &model_path,
1772 legacy_candle_device()?,
1773 dtype,
1774 )?;
1775 Ok(Arc::new(executor))
1776 }
1777 _ => Err(FerrumError::model(format!(
1778 "Architecture {:?} not supported",
1779 model_def.architecture
1780 ))),
1781 }
1782 }
1783
1784 fn metadata(&self) -> ComponentMetadata {
1785 ComponentMetadata {
1786 name: "llm".to_string(),
1787 version: "0.2.0".to_string(),
1788 description: "LLM executor (LlamaFamily / Qwen3MoE via Backend<B>; \
1789 BERT / CLIP / Whisper via candle)"
1790 .to_string(),
1791 supported_devices: cpu_cuda_and_optional_metal_devices(),
1792 capabilities: vec![
1793 "llama".to_string(),
1794 "qwen2".to_string(),
1795 "qwen3".to_string(),
1796 "qwen3_moe".to_string(),
1797 "mistral".to_string(),
1798 "bert".to_string(),
1799 "clip".to_string(),
1800 "whisper".to_string(),
1801 "safetensors".to_string(),
1802 "gguf".to_string(),
1803 "fp16".to_string(),
1804 "fp32".to_string(),
1805 ],
1806 }
1807 }
1808}
1809
1810pub struct StubTokenizer {
1816 vocab_size: usize,
1817 info: ferrum_interfaces::TokenizerInfo,
1818}
1819
1820impl StubTokenizer {
1821 pub fn new(vocab_size: usize) -> Self {
1823 let info = ferrum_interfaces::TokenizerInfo {
1824 tokenizer_type: ferrum_interfaces::tokenizer::TokenizerType::BPE,
1825 vocab_size,
1826 special_tokens: ferrum_types::SpecialTokens::default(),
1827 supports_incremental: false,
1828 supports_chat_template: false,
1829 max_token_length: None,
1830 model_name: Some("stub".into()),
1831 };
1832
1833 Self { vocab_size, info }
1834 }
1835}
1836
1837impl std::fmt::Debug for StubTokenizer {
1838 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1839 f.debug_struct("StubTokenizer")
1840 .field("vocab_size", &self.vocab_size)
1841 .finish()
1842 }
1843}
1844
1845impl Tokenizer for StubTokenizer {
1846 fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<ferrum_types::TokenId>> {
1847 let tokens: Vec<ferrum_types::TokenId> = text
1848 .split_whitespace()
1849 .enumerate()
1850 .map(|(i, _)| ferrum_types::TokenId::new((i % self.vocab_size) as u32))
1851 .collect();
1852
1853 Ok(if tokens.is_empty() {
1854 vec![ferrum_types::TokenId::new(0)]
1855 } else {
1856 tokens
1857 })
1858 }
1859
1860 fn decode(&self, tokens: &[ferrum_types::TokenId], _skip_special: bool) -> Result<String> {
1861 Ok(tokens
1862 .iter()
1863 .map(|t| format!("token_{}", t.get()))
1864 .collect::<Vec<_>>()
1865 .join(" "))
1866 }
1867
1868 fn decode_incremental(
1869 &self,
1870 _prev: &[ferrum_types::TokenId],
1871 next: ferrum_types::TokenId,
1872 ) -> Result<String> {
1873 Ok(format!("token_{} ", next.get()))
1874 }
1875
1876 fn vocab_size(&self) -> usize {
1877 self.vocab_size
1878 }
1879
1880 fn special_tokens(&self) -> &ferrum_types::SpecialTokens {
1881 &self.info.special_tokens
1882 }
1883
1884 fn token_id(&self, _text: &str) -> Option<ferrum_types::TokenId> {
1885 Some(ferrum_types::TokenId::new(0))
1886 }
1887
1888 fn token_text(&self, _token_id: ferrum_types::TokenId) -> Option<&str> {
1889 None
1890 }
1891
1892 fn info(&self) -> ferrum_interfaces::TokenizerInfo {
1893 self.info.clone()
1894 }
1895}
1896
1897static GLOBAL_REGISTRY: OnceLock<Arc<ComponentRegistry>> = OnceLock::new();
1903
1904pub fn global_registry() -> Arc<ComponentRegistry> {
1906 GLOBAL_REGISTRY
1907 .get_or_init(|| {
1908 info!("Initializing global component registry");
1909 Arc::new(ComponentRegistry::with_defaults())
1910 })
1911 .clone()
1912}
1913
1914pub fn set_global_registry(registry: Arc<ComponentRegistry>) -> Result<()> {
1916 GLOBAL_REGISTRY
1917 .set(registry)
1918 .map_err(|_| FerrumError::internal("Global registry already initialized"))
1919}
1920
1921#[cfg(test)]
1926mod tests {
1927 use super::*;
1928 use safetensors::tensor::{serialize_to_file, Dtype, TensorView};
1929 use std::path::{Path, PathBuf};
1930
1931 fn unique_test_dir(name: &str) -> PathBuf {
1932 let mut dir = std::env::temp_dir();
1933 dir.push(format!(
1934 "ferrum-{name}-{}-{}",
1935 std::process::id(),
1936 std::time::SystemTime::now()
1937 .duration_since(std::time::UNIX_EPOCH)
1938 .unwrap()
1939 .as_nanos()
1940 ));
1941 std::fs::create_dir_all(&dir).unwrap();
1942 dir
1943 }
1944
1945 #[test]
1946 fn fixed_storage_loaders_preserve_default_and_reject_other_kv_dtypes() {
1947 use ferrum_types::KvCacheDtype;
1948 for loader in ["legacy GGUF loader", "legacy Candle executor"] {
1949 validate_fixed_storage_kv_dtype(KvCacheDtype::Fp16, loader).unwrap();
1950 for dtype in [KvCacheDtype::Int8, KvCacheDtype::Bf16, KvCacheDtype::Fp8] {
1951 let error = validate_fixed_storage_kv_dtype(dtype, loader).unwrap_err();
1952 assert!(error.to_string().contains(loader));
1953 assert!(error.to_string().contains(dtype.as_str()));
1954 }
1955 }
1956 }
1957
1958 #[test]
1959 fn fixed_storage_factory_rejects_int8_before_loading_legacy_weights() {
1960 let directory = unique_test_dir("fixed-storage-kv-rejection");
1961 let gguf = directory.join("unloaded.gguf");
1962 std::fs::write(&gguf, b"not model weights").unwrap();
1965 let candle = directory.join("candle");
1966 std::fs::create_dir(&candle).unwrap();
1967 std::fs::write(
1968 candle.join("config.json"),
1969 serde_json::to_vec(&serde_json::json!({
1970 "architectures": ["BertModel"],
1971 "model_type": "bert",
1972 "hidden_size": 4,
1973 "num_hidden_layers": 1,
1974 "num_attention_heads": 1,
1975 "intermediate_size": 4,
1976 "vocab_size": 8,
1977 "max_position_embeddings": 8,
1978 }))
1979 .unwrap(),
1980 )
1981 .unwrap();
1982 for (path, loader) in [
1983 (&gguf, "legacy GGUF loader"),
1984 (&candle, "legacy Candle executor"),
1985 ] {
1986 let mut engine = EngineConfig::default();
1987 engine.kv_cache.dtype = ferrum_types::KvCacheDtype::Int8;
1988 engine.backend.device = Device::CPU;
1989 engine.backend.backend_options.insert(
1990 "model_path".to_owned(),
1991 serde_json::Value::String(path.to_string_lossy().into_owned()),
1992 );
1993 let component = ComponentConfig::from_engine_config(&engine);
1994 let error = match tokio_test::block_on(LlmExecutorFactory.create(&component)) {
1995 Ok(_) => panic!("a fixed-storage loader accepted INT8 KV"),
1996 Err(error) => error,
1997 };
1998 assert!(
1999 error
2000 .to_string()
2001 .contains(&format!("{loader} does not support KV dtype int8")),
2002 "unexpected early error: {error}"
2003 );
2004 }
2005 std::fs::remove_dir_all(directory).unwrap();
2006 }
2007
2008 fn write_qwen35_fixture_config(dir: &Path) {
2009 let config = serde_json::json!({
2010 "architectures": ["Qwen3_5ForConditionalGeneration"],
2011 "model_type": "qwen3_5",
2012 "vocab_size": 3,
2013 "max_position_embeddings": 16,
2014 "rms_norm_eps": 1e-6,
2015 "rope_theta": 10000.0,
2016 "tie_word_embeddings": false,
2017 "text_config": {
2018 "model_type": "qwen3_5_text",
2019 "hidden_size": 2,
2020 "intermediate_size": 2,
2021 "num_hidden_layers": 2,
2022 "layer_types": ["linear_attention", "full_attention"],
2023 "linear_num_key_heads": 1,
2024 "linear_num_value_heads": 1,
2025 "linear_key_head_dim": 1,
2026 "linear_value_head_dim": 1,
2027 "linear_conv_kernel_dim": 2,
2028 "mamba_ssm_dtype": "float32",
2029 "head_dim": 2,
2030 "num_attention_heads": 1,
2031 "num_key_value_heads": 1,
2032 "vocab_size": 3,
2033 "max_position_embeddings": 16,
2034 "tie_word_embeddings": false
2035 }
2036 });
2037 std::fs::write(
2038 dir.join("config.json"),
2039 serde_json::to_string_pretty(&config).unwrap(),
2040 )
2041 .unwrap();
2042 }
2043
2044 fn write_qwen35_fixture_safetensors(dir: &Path) {
2045 let tensors: Vec<(String, Vec<f32>)> = vec![
2046 (
2047 "model.embed_tokens.weight".to_string(),
2048 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
2049 ),
2050 ("model.norm.weight".to_string(), vec![0.0, 0.0]),
2051 (
2052 "model.lm_head.weight".to_string(),
2053 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
2054 ),
2055 (
2056 "model.layers.0.input_layernorm.weight".to_string(),
2057 vec![0.0, 0.0],
2058 ),
2059 (
2060 "model.layers.0.post_attention_layernorm.weight".to_string(),
2061 vec![0.0, 0.0],
2062 ),
2063 (
2064 "model.layers.0.linear_attn.in_proj_qkv.weight".to_string(),
2065 vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
2066 ),
2067 (
2068 "model.layers.0.linear_attn.in_proj_z.weight".to_string(),
2069 vec![1.0, -1.0],
2070 ),
2071 (
2072 "model.layers.0.linear_attn.in_proj_b.weight".to_string(),
2073 vec![0.5, 0.25],
2074 ),
2075 (
2076 "model.layers.0.linear_attn.in_proj_a.weight".to_string(),
2077 vec![-0.25, 0.75],
2078 ),
2079 (
2080 "model.layers.0.linear_attn.conv1d.weight".to_string(),
2081 vec![0.0, 1.0, 0.0, 1.0, 0.0, 1.0],
2082 ),
2083 ("model.layers.0.linear_attn.A_log".to_string(), vec![0.0]),
2084 ("model.layers.0.linear_attn.dt_bias".to_string(), vec![0.0]),
2085 (
2086 "model.layers.0.linear_attn.norm.weight".to_string(),
2087 vec![1.0],
2088 ),
2089 (
2090 "model.layers.0.linear_attn.out_proj.weight".to_string(),
2091 vec![1.0, -0.5],
2092 ),
2093 (
2094 "model.layers.0.mlp.gate_proj.weight".to_string(),
2095 vec![0.2, 0.1, -0.1, 0.3],
2096 ),
2097 (
2098 "model.layers.0.mlp.up_proj.weight".to_string(),
2099 vec![0.4, -0.2, 0.3, 0.5],
2100 ),
2101 (
2102 "model.layers.0.mlp.down_proj.weight".to_string(),
2103 vec![1.0, 0.0, 0.0, 1.0],
2104 ),
2105 (
2106 "model.layers.1.input_layernorm.weight".to_string(),
2107 vec![0.0, 0.0],
2108 ),
2109 (
2110 "model.layers.1.post_attention_layernorm.weight".to_string(),
2111 vec![0.0, 0.0],
2112 ),
2113 (
2114 "model.layers.1.self_attn.q_proj.weight".to_string(),
2115 vec![1.0, 0.0, 0.0, 1.0],
2116 ),
2117 (
2118 "model.layers.1.self_attn.k_proj.weight".to_string(),
2119 vec![0.5, 0.0, 0.0, 0.5],
2120 ),
2121 (
2122 "model.layers.1.self_attn.v_proj.weight".to_string(),
2123 vec![1.0, 1.0, -0.5, 0.5],
2124 ),
2125 (
2126 "model.layers.1.self_attn.o_proj.weight".to_string(),
2127 vec![1.0, 0.0, 0.0, 1.0],
2128 ),
2129 (
2130 "model.layers.1.self_attn.q_norm.weight".to_string(),
2131 vec![1.0, 1.0],
2132 ),
2133 (
2134 "model.layers.1.self_attn.k_norm.weight".to_string(),
2135 vec![1.0, 1.0],
2136 ),
2137 (
2138 "model.layers.1.mlp.gate_proj.weight".to_string(),
2139 vec![-0.2, 0.2, 0.1, 0.3],
2140 ),
2141 (
2142 "model.layers.1.mlp.up_proj.weight".to_string(),
2143 vec![0.25, 0.5, -0.3, 0.4],
2144 ),
2145 (
2146 "model.layers.1.mlp.down_proj.weight".to_string(),
2147 vec![0.5, 0.25, -0.2, 0.75],
2148 ),
2149 ];
2150 let views = tensors
2151 .into_iter()
2152 .map(|(name, values)| {
2153 let dimensions =
2154 if name.ends_with("embed_tokens.weight") || name.ends_with("lm_head.weight") {
2155 vec![3, 2]
2156 } else if name.ends_with("in_proj_qkv.weight") {
2157 vec![3, 2]
2158 } else if ["in_proj_z.weight", "in_proj_b.weight", "in_proj_a.weight"]
2159 .iter()
2160 .any(|suffix| name.ends_with(suffix))
2161 {
2162 vec![1, 2]
2163 } else if name.ends_with("linear_attn.out_proj.weight") {
2164 vec![2, 1]
2165 } else if name.ends_with("conv1d.weight") {
2166 vec![3, 1, 2]
2167 } else if values.len() == 4 {
2168 vec![2, 2]
2169 } else {
2170 vec![values.len()]
2171 };
2172 let bytes = values
2173 .iter()
2174 .flat_map(|value| value.to_le_bytes())
2175 .collect::<Vec<_>>()
2176 .into_boxed_slice();
2177 let bytes: &'static [u8] = Box::leak(bytes);
2178 (
2179 name,
2180 TensorView::new(Dtype::F32, dimensions, bytes).unwrap(),
2181 )
2182 })
2183 .collect::<Vec<_>>();
2184 serialize_to_file(
2185 views,
2186 &None::<std::collections::HashMap<String, String>>,
2187 &dir.join("model.safetensors"),
2188 )
2189 .unwrap();
2190 }
2191
2192 fn write_qwen35_fixture_model_dir() -> PathBuf {
2193 let dir = unique_test_dir("qwen35-vnext-route");
2194 write_qwen35_fixture_config(&dir);
2195 write_qwen35_fixture_safetensors(&dir);
2196 dir
2197 }
2198
2199 fn qwen35_fixture_component_config(model_dir: &Path) -> ComponentConfig {
2200 let mut engine_config = EngineConfig::default();
2201 engine_config.backend.device = Device::CPU;
2202 engine_config.backend.backend_options.insert(
2203 "model_path".to_string(),
2204 serde_json::Value::String(model_dir.to_string_lossy().to_string()),
2205 );
2206 ComponentConfig::from_engine_config(&engine_config)
2207 }
2208
2209 #[tokio::test]
2210 async fn startup_memory_real_cpu_plan_fits_before_upload_and_reaches_engine_config() {
2211 use ferrum_interfaces::vnext::{
2212 AllocationLifetime, DeviceId, DynamicResourceDemand, WeightMaterializerSelection,
2213 };
2214 use ferrum_kernels::backend::cpu::vnext_ops::CpuVNextComposition;
2215 use ferrum_types::{DeviceMemorySnapshot, StartupMemoryRequest};
2216
2217 let dir = write_qwen35_fixture_model_dir();
2218 let mut metadata: serde_json::Value =
2219 serde_json::from_slice(&std::fs::read(dir.join("config.json")).unwrap()).unwrap();
2220 let declared_context = 1 << 20;
2221 metadata["max_position_embeddings"] = serde_json::json!(declared_context);
2222 metadata["text_config"]["max_position_embeddings"] = serde_json::json!(declared_context);
2223 std::fs::write(
2224 dir.join("config.json"),
2225 serde_json::to_vec(&metadata).unwrap(),
2226 )
2227 .unwrap();
2228 std::fs::write(dir.join("tokenizer.json"), br#"{"version":"1.0","truncation":null,"padding":null,"added_tokens":[],"normalizer":null,"pre_tokenizer":{"type":"Whitespace"},"post_processor":null,"decoder":null,"model":{"type":"WordLevel","vocab":{"<unk>":0,"hello":1,"<eos>":2},"unk_token":"<unk>"}}"#).unwrap();
2229 std::fs::write(dir.join("tokenizer_config.json"), br#"{"chat_template":"{% for message in messages %}{{ message['content'] }}{% endfor %}","eos_token_id":2,"unk_token":"<unk>"}"#).unwrap();
2230 let defined = Arc::new(ferrum_models::vnext::qwen35::define_from_model_dir(&dir).unwrap());
2231 let capacity = 16 * 1024 * 1024;
2232 let mut config = qwen35_fixture_component_config(&dir).engine_config;
2233 config.scheduler.max_running_requests = 4;
2234 config.batching.max_num_batched_tokens = 8;
2235 config.backend.enable_reusable_execution = false;
2236 let create = |engine: &EngineConfig| {
2237 let composition = CpuVNextComposition::create(
2238 DeviceId::new(format!("device.cpu.startup-{}", uuid::Uuid::new_v4())).unwrap(),
2239 capacity,
2240 )
2241 .unwrap();
2242 let (runtime, operations, materializers, materializer, catalog) =
2243 composition.into_parts().unwrap();
2244 let observed = Arc::clone(&runtime);
2245 let executor = crate::product_composition::create_vnext_executor(
2246 engine,
2247 &defined,
2248 runtime,
2249 operations,
2250 materializers,
2251 catalog,
2252 |_| Ok(WeightMaterializerSelection::exact(materializer.clone())),
2253 );
2254 (observed, executor)
2255 };
2256 let (_, baseline) = create(&config);
2257 let baseline = baseline.unwrap();
2258 let memory = baseline
2259 .resolved_model_plan()
2260 .unwrap()
2261 .execution_plan()
2262 .payload()
2263 .memory();
2264 let tokens_per_physical_quantum = memory
2265 .dynamic_descriptors()
2266 .iter()
2267 .filter(|descriptor| descriptor.lifetime() == AllocationLifetime::Sequence)
2268 .filter_map(|descriptor| match descriptor.demand() {
2269 DynamicResourceDemand::Tokens {
2270 bytes_per_token, ..
2271 } => Some(
2272 descriptor
2273 .physical_allocation_quantum_bytes()
2274 .div_ceil(*bytes_per_token),
2275 ),
2276 _ => None,
2277 })
2278 .max()
2279 .expect("the hybrid fixture declares token-scaled attention state");
2280 let fitted_context_floor = tokens_per_physical_quantum * 2;
2281 let budget = memory
2282 .startup_peak_bytes(fitted_context_floor, 1, 8)
2283 .unwrap()
2284 .max(
2285 memory
2286 .startup_workload_peak_bytes(fitted_context_floor, 1, 2, 2)
2287 .unwrap(),
2288 );
2289 assert!(memory.startup_peak_bytes(declared_context, 1, 8).unwrap() > budget);
2290 drop(baseline);
2291 config.runtime.startup_memory_request = Some(StartupMemoryRequest {
2292 device: DeviceMemorySnapshot {
2293 capacity_bytes: capacity,
2294 available_bytes: capacity,
2295 source: "CPU fixture capacity".into(),
2296 },
2297 usable_capacity_bytes: budget,
2298 context_is_explicit: false,
2299 sequences_is_explicit: false,
2300 batch_is_explicit: false,
2301 });
2302 let (runtime, executor) = create(&config);
2303 let executor = executor.unwrap();
2304 let report = executor.startup_memory_plan().unwrap().clone();
2305 assert!(
2306 report.selected.context_tokens >= fitted_context_floor as usize
2307 && report.selected.context_tokens < declared_context as usize
2308 );
2309 assert!(report.selected.max_sequences > 1);
2310 assert_eq!(executor.kv_capacity(), Some(report.selected.context_tokens));
2311 let memory = executor
2312 .resolved_model_plan()
2313 .unwrap()
2314 .execution_plan()
2315 .payload()
2316 .memory();
2317 assert_eq!(
2318 report.context_peak_bytes,
2319 memory
2320 .startup_peak_bytes(
2321 report.selected.context_tokens as u64,
2322 1,
2323 report
2324 .selected
2325 .max_batch_tokens
2326 .min(report.selected.context_tokens) as u64,
2327 )
2328 .unwrap()
2329 );
2330 assert_eq!(
2331 report.decode_peak_bytes,
2332 memory
2333 .startup_workload_peak_bytes(
2334 report.selected.context_tokens as u64,
2335 1,
2336 report.selected.max_sequences as u32,
2337 report.selected.max_sequences as u64,
2338 )
2339 .unwrap()
2340 );
2341 assert!(runtime.peak_resident_bytes() > 0);
2342 let engine = crate::EngineBuilder::new(config.clone())
2343 .with_defined_model(Arc::clone(&defined))
2344 .with_custom_executor(Arc::new(executor))
2345 .build()
2346 .await
2347 .unwrap();
2348 assert_eq!(
2349 engine.config().runtime.startup_memory_plan.as_ref(),
2350 Some(&report)
2351 );
2352 assert_eq!(
2353 engine.context_capacity(),
2354 Some(report.selected.context_tokens)
2355 );
2356 assert_eq!(
2357 engine.config().scheduler.max_running_requests,
2358 report.selected.max_sequences
2359 );
2360 assert_eq!(
2361 engine.config().batching.max_num_batched_tokens,
2362 report.selected.max_batch_tokens
2363 );
2364 engine.shutdown().await.unwrap();
2365
2366 config.runtime.max_model_len = Some(declared_context as usize);
2367 let request = config.runtime.startup_memory_request.as_mut().unwrap();
2368 request.context_is_explicit = true;
2369 request.batch_is_explicit = true;
2370 let (runtime, rejected) = create(&config);
2371 let error = rejected
2372 .err()
2373 .expect("explicit capacity must not silently shrink");
2374 assert!(error.to_string().contains("explicit"), "{error}");
2375 assert_eq!(
2376 runtime.peak_resident_bytes(),
2377 0,
2378 "capacity rejection must precede weight allocation"
2379 );
2380 std::fs::remove_dir_all(dir).unwrap();
2381 }
2382
2383 fn enable_checkpoint_capture(config: &mut ComponentConfig, output_dir: PathBuf) {
2384 config.engine_config.runtime.vnext_checkpoint_capture =
2385 Some(ferrum_types::VNextCheckpointCaptureConfig {
2386 output_dir,
2387 value_ids: vec!["value.output.logits".to_owned()],
2388 maximum_prefill_waves: 1,
2389 maximum_decode_waves: 0,
2390 capture_product_output: false,
2391 teacher_forcing: None,
2392 });
2393 }
2394
2395 #[test]
2396 fn checkpoint_capture_rejects_stub_fallback() {
2397 let mut config = ComponentConfig::from_engine_config(&EngineConfig::default());
2398 enable_checkpoint_capture(&mut config, PathBuf::from("capture"));
2399
2400 let error = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2401 Ok(_) => panic!("checkpoint capture unexpectedly entered the stub executor"),
2402 Err(error) => error.to_string(),
2403 };
2404 assert!(
2405 error.contains("requires a registered model source"),
2406 "{error}"
2407 );
2408 }
2409
2410 #[test]
2411 fn qwen35_cpu_registry_requires_typed_sources_before_legacy_loading() {
2412 let dir = write_qwen35_fixture_model_dir();
2413 std::fs::remove_file(dir.join("model.safetensors")).unwrap();
2414 let config = qwen35_fixture_component_config(&dir);
2415
2416 let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2417 Ok(_) => panic!("registered Qwen3.5 package accepted missing tokenizer and weights"),
2418 Err(err) => err.to_string(),
2419 };
2420
2421 assert!(err.contains("tokenizer.json"), "{err}");
2422 validate_registered_vnext_backend(
2423 ferrum_models::vnext::ProductionExecutionKind::CausalLanguage,
2424 &Device::CPU,
2425 &ferrum_interfaces::vnext::ExternalModelMetadataId::new(
2426 ferrum_models::vnext::qwen35::EXTERNAL_METADATA_ID,
2427 )
2428 .unwrap(),
2429 )
2430 .unwrap();
2431 let _ = std::fs::remove_dir_all(dir);
2432 }
2433
2434 #[test]
2435 fn qwen35_typed_gguf_routes_registered_package_before_legacy_gguf_loading() {
2436 let dir = unique_test_dir("qwen35-typed-gguf-route");
2437 write_qwen35_fixture_config(&dir);
2438 std::fs::write(dir.join("tokenizer.json"), br#"{"version":"1.0"}"#).unwrap();
2439 std::fs::write(
2440 dir.join("tokenizer_config.json"),
2441 br#"{"chat_template":"fixture","eos_token_id":2}"#,
2442 )
2443 .unwrap();
2444 let gguf = dir.join("model.gguf");
2445 let mut bytes = b"GGUF".to_vec();
2449 bytes.extend_from_slice(&3_u32.to_le_bytes());
2450 bytes.extend_from_slice(&1_u64.to_le_bytes()); bytes.extend_from_slice(&1_u64.to_le_bytes()); let key = "general.architecture";
2453 bytes.extend_from_slice(&(key.len() as u64).to_le_bytes());
2454 bytes.extend_from_slice(key.as_bytes());
2455 bytes.extend_from_slice(&8_u32.to_le_bytes()); bytes.extend_from_slice(&6_u64.to_le_bytes());
2457 bytes.extend_from_slice(b"qwen35");
2458 let name = "output_norm.weight";
2459 bytes.extend_from_slice(&(name.len() as u64).to_le_bytes());
2460 bytes.extend_from_slice(name.as_bytes());
2461 bytes.extend_from_slice(&1_u32.to_le_bytes()); bytes.extend_from_slice(&2_u64.to_le_bytes()); bytes.extend_from_slice(&0_u32.to_le_bytes()); bytes.extend_from_slice(&0_u64.to_le_bytes()); bytes.resize(bytes.len().div_ceil(32) * 32, 0);
2466 bytes.extend_from_slice(&1_f32.to_le_bytes());
2467 bytes.extend_from_slice(&1_f32.to_le_bytes());
2468 std::fs::write(&gguf, bytes).unwrap();
2469 let source = ferrum_quantization::gguf::NativeGgufFile::open(&gguf).unwrap();
2470 assert_eq!(source.architecture().unwrap(), "qwen35");
2471 assert_eq!(source.tensor_count(), 1);
2472 assert_eq!(source.tensor_info(name).unwrap().dimensions, [2]);
2473 drop(source);
2474 let original = ferrum_interfaces::vnext::OriginalModelSource {
2475 kind: ferrum_interfaces::vnext::ModelSourceKind::LocalDirectory,
2476 location: dir.display().to_string(),
2477 requested_revision: None,
2478 };
2479 let sources = Arc::new(
2480 ProductionModelSourceBundle::open(
2481 &dir,
2482 &dir,
2483 ferrum_models::vnext::ProductionWeightArtifact::gguf_file(&gguf),
2484 ferrum_interfaces::vnext::OriginalModelSources {
2485 semantic: original.clone(),
2486 tokenizer: original.clone(),
2487 weights: ferrum_interfaces::vnext::OriginalModelSource {
2488 kind: ferrum_interfaces::vnext::ModelSourceKind::LocalFile,
2489 location: gguf.display().to_string(),
2490 requested_revision: None,
2491 },
2492 },
2493 )
2494 .unwrap(),
2495 );
2496 let mut engine = EngineConfig::default();
2497 engine.backend.device = Device::CPU;
2498 let mut config = ComponentConfig::from_engine_config(&engine);
2499 config.model_sources = Some(sources);
2500
2501 let err = match tokio_test::block_on(LlmExecutorFactory.create(&config)) {
2502 Ok(_) => panic!("registered Qwen3.5 CPU composition accepted missing model tensors"),
2503 Err(err) => err.to_string(),
2504 };
2505
2506 assert!(
2507 err.contains("Qwen3.5 GGUF is missing required role \"embed_tokens\""),
2508 "{err}"
2509 );
2510 assert!(err.contains("token_embd.weight"), "{err}");
2511 let _ = std::fs::remove_dir_all(dir);
2512 }
2513
2514 #[test]
2515 fn test_registry_creation() {
2516 let registry = ComponentRegistry::new();
2517 assert!(registry.list_tokenizers().is_empty());
2518 }
2519
2520 #[test]
2521 fn test_registry_with_defaults() {
2522 let registry = ComponentRegistry::with_defaults();
2523 assert!(registry.list_tokenizers().contains(&"stub".to_string()));
2524 assert!(registry
2525 .list_samplers()
2526 .contains(&"multinomial".to_string()));
2527 assert!(registry.list_schedulers().contains(&"fifo".to_string()));
2528 }
2529
2530 #[test]
2531 fn test_component_config() {
2532 let mut options = HashMap::new();
2533 options.insert("test_key".to_string(), serde_json::json!("test_value"));
2534
2535 let config = ComponentConfig {
2536 engine_config: EngineConfig::default(),
2537 device: Device::CPU,
2538 component_options: options,
2539 model_sources: None,
2540 defined_model: None,
2541 };
2542
2543 assert_eq!(
2544 config.get_string_option("test_key"),
2545 Some("test_value".to_string())
2546 );
2547 assert!(config.get_string_option("missing").is_none());
2548 }
2549
2550 fn layer_split_component_config_for_device(device: usize) -> ComponentConfig {
2551 let mut config = EngineConfig::default();
2552 config.backend.device = Device::CUDA(device);
2553 config.backend.backend_options.insert(
2554 "selected_distributed_strategy".to_string(),
2555 serde_json::Value::String("layer_split".to_string()),
2556 );
2557 config.backend.backend_options.insert(
2558 "selected_gpu_devices".to_string(),
2559 serde_json::json!([0, 1]),
2560 );
2561 config.backend.backend_options.insert(
2562 "selected_layer_split_plan".to_string(),
2563 serde_json::Value::String(
2564 "stage0:cuda:0:layers=0-39;stage1:cuda:1:layers=40-79".to_string(),
2565 ),
2566 );
2567 config.backend.backend_options.insert(
2568 "selected_layer_split_stages".to_string(),
2569 serde_json::json!([
2570 {"stage": 0, "device": 0, "layer_start": 0, "layer_end": 39},
2571 {"stage": 1, "device": 1, "layer_start": 40, "layer_end": 79}
2572 ]),
2573 );
2574 ComponentConfig::from_engine_config(&config)
2575 }
2576
2577 #[test]
2578 fn resolves_llama_stage_config_for_selected_cuda_device() {
2579 let config = layer_split_component_config_for_device(1);
2580
2581 let stage = resolve_llama_layer_stage_config(&config, 80)
2582 .unwrap()
2583 .unwrap();
2584
2585 assert_eq!(stage.source_layers, 40..80);
2586 assert!(!stage.load_embedding);
2587 assert!(stage.load_lm_head);
2588 }
2589
2590 #[test]
2591 fn resolves_layer_split_pipeline_mode_default_and_explicit_batch() {
2592 let mut config = layer_split_component_config_for_device(0);
2593
2594 assert_eq!(
2595 resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2596 ferrum_models::models::LlamaPipelineMode::Overlapped
2597 );
2598
2599 config.component_options.insert(
2600 "layer_split_pipeline_mode".to_string(),
2601 serde_json::Value::String("batch".to_string()),
2602 );
2603 assert_eq!(
2604 resolve_llama_layer_split_pipeline_mode(&config, 2).unwrap(),
2605 ferrum_models::models::LlamaPipelineMode::Batch
2606 );
2607 }
2608
2609 #[test]
2610 fn rejects_llama_stage_config_when_plan_layer_count_mismatches_model() {
2611 let config = layer_split_component_config_for_device(0);
2612
2613 let err = resolve_llama_layer_stage_config(&config, 81)
2614 .unwrap_err()
2615 .to_string();
2616
2617 assert!(err.contains("covers 80 layers but model has 81"));
2618 }
2619
2620 #[test]
2621 fn build_llm_rejects_layer_split_before_weight_open_without_device_scope() {
2622 use ferrum_interfaces::kv_dtype::KvFp16;
2623 use ferrum_models::models::LlamaFamilyConfig;
2624
2625 let qcfg = LlamaFamilyConfig {
2626 hidden_size: 1,
2627 intermediate_size: 1,
2628 num_heads: 1,
2629 num_kv_heads: 1,
2630 head_dim: 1,
2631 num_layers: 2,
2632 vocab_size: 1,
2633 max_seq_len: 8,
2634 rms_norm_eps: 1e-5,
2635 rope_theta: 10_000.0,
2636 rope_scaling: None,
2637 rope_interleaved: false,
2638 has_qk_norm: false,
2639 sliding_window: 0,
2640 ..Default::default()
2641 };
2642 let plan = crate::layer_split::parse_layer_split_plan(
2643 "stage0:cuda:0:layers=0-0;stage1:cuda:1:layers=1-1",
2644 )
2645 .unwrap();
2646
2647 let err = match build_llm::<ferrum_kernels::backend::cpu::CpuBackend, KvFp16>(
2648 ferrum_models::Architecture::Llama,
2649 qcfg,
2650 None,
2651 "/missing/model/path",
2652 Some(plan),
2653 Some(ferrum_models::models::LlamaPipelineMode::Overlapped),
2654 ) {
2655 Ok(_) => panic!("layer_split build unexpectedly succeeded"),
2656 Err(err) => err.to_string(),
2657 };
2658
2659 assert!(err.contains("device-scoped execution"));
2660 assert!(err.contains("refusing to silently use the default device"));
2661 }
2662
2663 #[test]
2664 fn test_registry_runtime_env_parses_model_dtype_and_tp() {
2665 let env = RegistryRuntimeEnv::from_env_vars([
2666 ("FERRUM_MODEL_PATH", "/models/qwen"),
2667 ("FERRUM_METAL_DTYPE", "fp16"),
2668 ("FERRUM_DTYPE", "fp32"),
2669 ("FERRUM_TP", "4"),
2670 ]);
2671
2672 assert_eq!(env.model_path.as_deref(), Some("/models/qwen"));
2673 assert_eq!(env.metal_dtype.as_deref(), Some("fp16"));
2674 assert_eq!(env.dtype.as_deref(), Some("fp32"));
2675 assert_eq!(env.tp, 4);
2676 }
2677
2678 #[test]
2679 fn test_registry_runtime_env_defaults_invalid_tp_and_cuda_dtype() {
2680 let env = RegistryRuntimeEnv::from_env_vars([
2681 ("FERRUM_DTYPE", "not-a-dtype"),
2682 ("FERRUM_TP", "not-a-number"),
2683 ]);
2684
2685 assert_eq!(env.model_path(), None);
2686 assert_eq!(env.tp, 0);
2687 assert_eq!(
2688 env.dtype_for_device(&Device::CUDA(0)),
2689 candle_core::DType::F16
2690 );
2691 assert_eq!(env.dtype_for_device(&Device::CPU), candle_core::DType::F32);
2692 }
2693
2694 #[cfg(any(target_os = "macos", target_os = "ios"))]
2695 #[test]
2696 fn test_registry_runtime_env_metal_dtype_precedence() {
2697 let env = RegistryRuntimeEnv::from_env_vars([
2698 ("FERRUM_METAL_DTYPE", "fp16"),
2699 ("FERRUM_DTYPE", "fp32"),
2700 ]);
2701
2702 assert_eq!(
2703 env.dtype_for_device(&Device::Metal),
2704 candle_core::DType::F16
2705 );
2706 assert_eq!(env.metal_dtype_hint(), "fp16");
2707 }
2708
2709 #[test]
2710 fn test_stub_tokenizer() {
2711 let tokenizer = StubTokenizer::new(100);
2712 assert_eq!(tokenizer.vocab_size(), 100);
2713
2714 let tokens = tokenizer.encode("hello world", false).unwrap();
2715 assert!(!tokens.is_empty());
2716
2717 let text = tokenizer.decode(&tokens, false).unwrap();
2718 assert!(text.contains("token_"));
2719 }
2720
2721 #[test]
2722 fn test_greedy_sampler() {
2723 use rand::SeedableRng;
2724
2725 let sampler = GreedySampler;
2726 let mut rng = rand::rngs::StdRng::seed_from_u64(42);
2727
2728 let logits = vec![0.1, 0.5, 0.3, 0.9, 0.2];
2729 let token = sampler.sample(&logits, &mut rng).unwrap();
2730
2731 assert_eq!(token.get(), 3); }
2733}