Cesium2 / src /audio.py
MORPH-AI
feat: dynamic MoE expansion, multi-head CoT, plugin architecture, improved MoD
82f262a
Raw
History Blame Contribute Delete
4.99 kB
"""
AudioModule - ASR (Automatic Speech Recognition) + TTS (Text-to-Speech)
for MORPH-AI v6.
Lazy-loads Whisper for ASR and Coqui TTS / gTTS for speech synthesis.
Falls back to feature-only mode when models are unavailable.
"""
import io
import json
import os
import tempfile
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from architecture import MorphConfig
@dataclass
class AudioFacts:
duration: float = 0.0
sample_rate: int = 16000
transcription: str = ""
language: str = "en"
confidence: float = 0.0
embedding: Optional[torch.Tensor] = None
segments: List[Dict[str, Any]] = field(default_factory=list)
def to_text(self) -> str:
parts = [f"audio {self.duration:.1f}s {self.sample_rate}Hz"]
if self.transcription:
parts.append(f"transcription: {self.transcription}")
if self.language != "en":
parts.append(f"language: {self.language}")
return " | ".join(parts)
def to_dict(self) -> Dict[str, Any]:
return {
"duration": self.duration,
"sample_rate": self.sample_rate,
"transcription": self.transcription,
"language": self.language,
"confidence": self.confidence,
}
class AudioModule(nn.Module):
"""ASR + TTS module with lazy model loading and fallback."""
def __init__(self, config: MorphConfig, hidden_dim: int):
super().__init__()
self.hidden_dim = hidden_dim
self.audio_proj = nn.Linear(config.audio_dim, hidden_dim)
self.whisper = None
self.whisper_processor = None
self.tts_model = None
self._loaded = False
def _load_models(self, device: str = "cpu"):
if self._loaded:
return
try:
from transformers import WhisperForConditionalGeneration, WhisperProcessor
self.whisper = WhisperForConditionalGeneration.from_pretrained(
"openai/whisper-tiny"
).to(device).eval()
self.whisper_processor = WhisperProcessor.from_pretrained("openai/whisper-tiny")
print("Whisper ASR loaded")
except Exception as e:
print(f"Whisper load failed: {e}")
try:
from TTS.api import TTS
self.tts_model = TTS(model_name="tts_models/en/ljspeech/tacotron2-DDC", progress_bar=False)
print("Coqui TTS loaded")
except Exception as e:
print(f"TTS load failed: {e}")
self._loaded = True
def transcribe(self, audio_source, device: str = "cpu") -> AudioFacts:
"""Transcribe audio to text using Whisper ASR."""
self._load_models(device)
facts = AudioFacts()
try:
import librosa
audio, sr = librosa.load(audio_source, sr=16000)
facts.duration = librosa.get_duration(y=audio, sr=sr)
facts.sample_rate = sr
if self.whisper is not None and self.whisper_processor is not None:
inputs = self.whisper_processor(audio, sampling_rate=sr, return_tensors="pt").to(device)
with torch.no_grad():
generated = self.whisper.generate(inputs.input_features)
transcription = self.whisper_processor.batch_decode(generated, skip_special_tokens=True)[0]
facts.transcription = transcription
facts.confidence = 0.9
else:
facts.transcription = "[ASR unavailable - whisper not loaded]"
except ImportError:
facts.transcription = "[ASR requires librosa + transformers: pip install librosa transformers]"
except Exception as e:
facts.transcription = f"[ASR error: {e}]"
return facts
def synthesize(self, text: str, output_path: Optional[str] = None, device: str = "cpu") -> Optional[str]:
"""Synthesize speech from text using TTS."""
self._load_models(device)
if output_path is None:
output_path = tempfile.mktemp(suffix=".wav")
try:
if self.tts_model is not None:
self.tts_model.tts_to_file(text=text, file_path=output_path)
return output_path
else:
from gtts import gTTS
tts = gTTS(text=text, lang="en")
tts.save(output_path)
return output_path
except ImportError:
print("TTS unavailable - install gTTS or Coqui TTS")
return None
except Exception as e:
print(f"TTS error: {e}")
return None
def forward(self, hidden: torch.Tensor, audio_embeds: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Project audio embeddings into hidden space."""
if audio_embeds is None:
return hidden
return hidden + self.audio_proj(audio_embeds.to(hidden.dtype))