| import sys, os, re, torch, torchaudio, numpy as np |
| from typing import Dict, Any |
|
|
|
|
| class EndpointHandler: |
| def __init__(self, path=""): |
| if path and path not in sys.path: |
| sys.path.insert(0, path) |
|
|
| from configuration_moss_audio import MossAudioConfig |
| from processing_moss_audio import MossAudioProcessor |
| from modeling_moss_audio import MossAudioModel |
|
|
| print(f"[MOSS-8B] Loading (no time markers)...") |
| self.config = MossAudioConfig.from_pretrained(path) |
| self.processor = MossAudioProcessor.from_pretrained( |
| path, trust_remote_code=True, enable_time_marker=False |
| ) |
| self.model = MossAudioModel.from_pretrained( |
| path, config=self.config, torch_dtype=torch.bfloat16, |
| device_map="auto", trust_remote_code=True, |
| ) |
| self.model.eval() |
| print(f"[MOSS-8B] Loaded.") |
| self.sr = self.processor.config.mel_sr |
|
|
| def _load_audio(self, audio_input): |
| if isinstance(audio_input, str) and audio_input.startswith(("http://", "https://")): |
| import requests, tempfile |
| resp = requests.get(audio_input, timeout=120) |
| resp.raise_for_status() |
| with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp: |
| tmp.write(resp.content) |
| tmp_path = tmp.name |
| waveform, sr = torchaudio.load(tmp_path) |
| os.unlink(tmp_path) |
| elif isinstance(audio_input, str): |
| waveform, sr = torchaudio.load(audio_input) |
| else: |
| raise ValueError("audio must be URL or file path") |
|
|
| if waveform.shape[0] > 1: |
| waveform = waveform.mean(dim=0, keepdim=True) |
| if sr != self.sr: |
| waveform = torchaudio.transforms.Resample(sr, self.sr)(waveform) |
| return waveform |
|
|
| def _transcribe(self, raw_audio, prompt, gen_kwargs): |
| inputs = self.processor(text=prompt, audios=[raw_audio], return_tensors="pt") |
| inputs = inputs.to(self.model.device) |
| if inputs.get("audio_data") is not None: |
| inputs["audio_data"] = inputs["audio_data"].to(self.model.dtype) |
| inputs["audio_input_mask"] = inputs["input_ids"] == self.processor.audio_token_id |
| ids = self.model.generate(**inputs, **gen_kwargs) |
| text = self.processor.decode(ids[0, inputs["input_ids"].shape[1]:], skip_special_tokens=True).strip() |
| text = re.sub(r'<think>.*?</think>\s*\n?\s*', '', text, flags=re.DOTALL).strip() |
| return text |
|
|
| @torch.no_grad() |
| def __call__(self, data): |
| inputs = data.get("inputs", data) |
| try: |
| if isinstance(inputs, str): |
| audio_input, prompt = inputs, "转写这段日语音频,标注说话人。格式:说话人:内容" |
| elif isinstance(inputs, dict): |
| audio_input = inputs.get("audio", "") |
| prompt = inputs.get("prompt", inputs.get("text", "转写这段日语音频,标注说话人。格式:说话人:内容")) |
| else: |
| return {"error": f"Invalid input type: {type(inputs)}"} |
|
|
| gen_kwargs = dict( |
| max_new_tokens=int(data.get("max_new_tokens", inputs.get("max_new_tokens", 384))), |
| num_beams=int(data.get("num_beams", inputs.get("num_beams", 2))), |
| use_cache=True, |
| do_sample=bool(data.get("do_sample", inputs.get("do_sample", False))), |
| temperature=float(data.get("temperature", inputs.get("temperature", 0.5))), |
| top_p=float(data.get("top_p", inputs.get("top_p", 0.9))), |
| top_k=int(data.get("top_k", inputs.get("top_k", 50))), |
| repetition_penalty=float(data.get("repetition_penalty", inputs.get("repetition_penalty", 1.15))), |
| no_repeat_ngram_size=int(data.get("no_repeat_ngram_size", inputs.get("no_repeat_ngram_size", 3))), |
| length_penalty=float(data.get("length_penalty", inputs.get("length_penalty", 0.9))), |
| early_stopping=True, |
| ) |
|
|
| waveform = self._load_audio(audio_input) |
| total_sec = waveform.shape[1] / self.sr |
| print(f"[MOSS-8B] Audio: {total_sec:.1f}s (no time markers)") |
|
|
| |
| chunk_sec = 120 |
| overlap_sec = 10 |
| step_samples = int((chunk_sec - overlap_sec) * self.sr) |
| chunk_samples = int(chunk_sec * self.sr) |
| n_frames = waveform.shape[1] |
|
|
| chunks_list = [] |
| pos = 0 |
| while pos < n_frames: |
| end = min(pos + chunk_samples, n_frames) |
| chunks_list.append((pos, end)) |
| pos += step_samples |
| if end >= n_frames: |
| break |
|
|
| segments = [] |
| prev_speakers = "" |
|
|
| for idx, (st, en) in enumerate(chunks_list): |
| chunk = waveform[:, st:en] |
| if chunk.abs().max() < 0.005: |
| text = "" |
| else: |
| ctx_prompt = prompt |
| if prev_speakers and idx > 0: |
| ctx_prompt = prompt + f" 前段说话人:{prev_speakers}" |
|
|
| raw = chunk.squeeze(0).numpy().astype(np.float32) |
| text = self._transcribe(raw, ctx_prompt, gen_kwargs) |
|
|
| speakers = set(re.findall(r'([\u3040-\u30ff\u4e00-\u9fff]+)[::]', text)) |
| if speakers: |
| prev_speakers = "、".join(sorted(speakers)[:4]) |
|
|
| if torch.cuda.is_available(): |
| torch.cuda.empty_cache() |
|
|
| start_t = st / self.sr |
| end_t = en / self.sr |
| segments.append({"start": round(start_t, 2), "end": round(end_t, 2), "text": text}) |
| print(f"[Chunk {idx+1}/{len(chunks_list)}] ({len(text)}c) {text[:80]}") |
|
|
| full_text = " ".join(s["text"] for s in segments if s["text"]) |
|
|
| srt_lines = [] |
| idx_n = 1 |
| for seg in segments: |
| t = seg["text"].strip() |
| if not t: |
| continue |
| s, e = seg["start"], seg["end"] |
| srt_lines.append(str(idx_n)) |
| srt_lines.append(f"{int(s//3600):02d}:{int(s%3600//60):02d}:{s%60:06.3f} --> {int(e//3600):02d}:{int(e%3600//60):02d}:{e%60:06.3f}") |
| srt_lines.append(t) |
| srt_lines.append("") |
| idx_n += 1 |
|
|
| return { |
| "transcription": full_text, |
| "segments": segments, |
| "srt": "\n".join(srt_lines), |
| "num_chunks": len(chunks_list), |
| } |
|
|
| except Exception as e: |
| import traceback |
| return {"error": str(e), "traceback": traceback.format_exc()} |
|
|