File size: 4,988 Bytes
82f262a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""
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))