Skip to content
Open
Show file tree
Hide file tree
Changes from all 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