Skip to content
Merged
Show file tree
Hide file tree
Changes from 8 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/mobius/_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,7 @@
Qwen35VLTextModel,
QwenCausalLMModel,
SmolLM3CausalLMModel,
SortformerDiarizationModel,
WhisperForConditionalGeneration,
)
from mobius.models.bamba import BambaCausalLMModel
Expand Down Expand Up @@ -765,6 +766,7 @@ def _detect_fallback_registration(hf_config) -> ModelRegistration | None:
"wavlm": ModelRegistration(Wav2Vec2Model, task="audio-feature-extraction"),
"mms": ModelRegistration(Wav2Vec2ForCTCModel, task="ctc-asr", config_class=MMSConfig),
"fastconformer_rnnt": ModelRegistration(EncDecRNNTModel, task="fastconformer-rnnt"),
"sortformer": ModelRegistration(SortformerDiarizationModel, task="diarization"),
}


Expand Down
17 changes: 15 additions & 2 deletions src/mobius/integrations/nemo/_config_mapping.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,11 +20,13 @@
from typing import Any

from mobius._configs import ArchitectureConfig
from mobius._configs._base import BaseModelConfig

# NeMo ``target`` class path → mobius registry model_type.
NEMO_TARGET_TO_MODEL_TYPE: dict[str, str] = {
"nemo.collections.asr.models.rnnt_bpe_models.EncDecRNNTBPEModel": "fastconformer_rnnt",
"nemo.collections.asr.models.rnnt_models.EncDecRNNTModel": "fastconformer_rnnt",
"nemo.collections.asr.models.sortformer_diar_models.SortformerEncLabelModel": "sortformer",
}


Expand Down Expand Up @@ -85,11 +87,22 @@ def _validate_encoder(enc: dict[str, Any]) -> None:
)


def nemo_to_config(nemo_config: dict[str, Any]) -> ArchitectureConfig:
"""Build an :class:`ArchitectureConfig` from a NeMo ``model_config.yaml`` dict."""
def nemo_to_config(nemo_config: dict[str, Any]) -> BaseModelConfig:
"""Build a mobius config from a NeMo ``model_config.yaml`` dict.

Dispatches on the NeMo ``target`` class path: FastConformer-RNNT models
produce an :class:`ArchitectureConfig`; Sortformer diarization models
produce a :class:`SortformerConfig`.
"""
target = str(nemo_config.get("target", ""))
model_type = nemo_model_type(target)

if model_type == "sortformer":
# Imported lazily to avoid a models→integrations import cycle.
from mobius.models.sortformer import SortformerConfig

return SortformerConfig.from_nemo_yaml(nemo_config)

enc = nemo_config["encoder"]
dec = nemo_config["decoder"]
joint = nemo_config["joint"]
Expand Down
3 changes: 3 additions & 0 deletions src/mobius/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,8 @@
"Qwen3CausalLMModel",
"Qwen3NextCausalLMModel",
"SenseVoiceSmallModel",
"SortformerConfig",
"SortformerDiarizationModel",
"Qwen3TTSCodePredictorModel",
"Qwen3TTSCodecDecoderModel",
"Qwen3TTSCodecEncoderModel",
Expand Down Expand Up @@ -287,6 +289,7 @@
Qwen25VLVisionEncoderModel,
)
from mobius.models.sensevoice_small import SenseVoiceSmallModel
from mobius.models.sortformer import SortformerConfig, SortformerDiarizationModel
from mobius.models.smollm import SmolLM3CausalLMModel
from mobius.models.starcoder2 import StarCoder2CausalLMModel
from mobius.models.t5 import T5ForConditionalGeneration
Expand Down
Loading
Loading