Initial commit: Børnetime app scaffold
FastAPI + Celery + Svelte app for English→Danish video translation with voice cloning. Includes diarization, STT, translation, TTS, and audio mixing pipeline.
This commit is contained in:
@@ -0,0 +1,54 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
|
||||
from ..models.database import get_session
|
||||
from ..models.tables import Job, Video, JobStatus
|
||||
from ..models.schemas import JobOut, JobListOut, ProcessResponse
|
||||
from ..tasks.process import process_video
|
||||
from ..tasks.generate import generate_audio
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("", response_model=JobListOut)
|
||||
async def list_jobs(session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Job).order_by(Job.created_at.desc()))
|
||||
jobs = result.scalars().all()
|
||||
return JobListOut(jobs=[JobOut.model_validate(j) for j in jobs])
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=JobOut)
|
||||
async def get_job(job_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Job).where(Job.id == job_id))
|
||||
job = result.scalar_one_or_none()
|
||||
if not job:
|
||||
raise HTTPException(404, "Job not found")
|
||||
return JobOut.model_validate(job)
|
||||
|
||||
|
||||
@router.post("/{job_id}/cancel", response_model=ProcessResponse)
|
||||
async def cancel_job(job_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Job).where(Job.id == job_id))
|
||||
job = result.scalar_one_or_none()
|
||||
if not job:
|
||||
raise HTTPException(404, "Job not found")
|
||||
|
||||
if job.status not in (JobStatus.pending.value, JobStatus.running.value):
|
||||
raise HTTPException(400, f"Cannot cancel job in status '{job.status}'")
|
||||
|
||||
job.status = JobStatus.failed.value
|
||||
job.error_message = "Cancelled by user"
|
||||
await session.commit()
|
||||
|
||||
video_result = await session.execute(select(Video).where(Video.id == job.video_id))
|
||||
video = video_result.scalar_one_or_none()
|
||||
if video:
|
||||
if job.job_type == "process":
|
||||
video.status = "uploaded"
|
||||
elif job.job_type == "generate":
|
||||
video.status = "review"
|
||||
video.error_message = None
|
||||
await session.commit()
|
||||
|
||||
return ProcessResponse(job_id=job.id, message="Job cancelled")
|
||||
@@ -0,0 +1,218 @@
|
||||
import os
|
||||
import logging
|
||||
import aiofiles
|
||||
from fastapi import APIRouter, Depends, UploadFile, File, HTTPException
|
||||
from fastapi.responses import FileResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
|
||||
from ..models.database import get_session
|
||||
from ..models.tables import Video, Segment, Job, VideoStatus, JobStatus
|
||||
from ..models.schemas import (
|
||||
VideoOut,
|
||||
VideoListOut,
|
||||
SegmentOut,
|
||||
SegmentUpdate,
|
||||
TranscriptOut,
|
||||
ProcessResponse,
|
||||
)
|
||||
from ..services.video_utils import get_video_duration
|
||||
from ..config import settings
|
||||
from ..tasks.process import process_video
|
||||
from ..tasks.generate import generate_audio
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("", response_model=VideoListOut)
|
||||
async def list_videos(session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).order_by(Video.created_at.desc()))
|
||||
videos = result.scalars().all()
|
||||
return VideoListOut(videos=[VideoOut.model_validate(v) for v in videos])
|
||||
|
||||
|
||||
@router.get("/{video_id}", response_model=VideoOut)
|
||||
async def get_video(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).where(Video.id == video_id))
|
||||
video = result.scalar_one_or_none()
|
||||
if not video:
|
||||
raise HTTPException(404, "Video not found")
|
||||
return VideoOut.model_validate(video)
|
||||
|
||||
|
||||
@router.post("/upload", response_model=VideoOut)
|
||||
async def upload_video(
|
||||
file: UploadFile = File(...), session: AsyncSession = Depends(get_session)
|
||||
):
|
||||
if not file.filename:
|
||||
raise HTTPException(400, "No filename")
|
||||
|
||||
allowed = {".mp4", ".mkv", ".mov", ".avi", ".webm", ".m4v"}
|
||||
ext = os.path.splitext(file.filename)[1].lower()
|
||||
if ext not in allowed:
|
||||
raise HTTPException(
|
||||
400, f"Unsupported format: {ext}. Allowed: {', '.join(allowed)}"
|
||||
)
|
||||
|
||||
dest = os.path.join(settings.video_dir, file.filename)
|
||||
async with aiofiles.open(dest, "wb") as f:
|
||||
while chunk := await file.read(64 * 1024):
|
||||
await f.write(chunk)
|
||||
|
||||
duration = get_video_duration(dest)
|
||||
|
||||
video = Video(
|
||||
filename=file.filename,
|
||||
filepath=dest,
|
||||
duration=duration,
|
||||
status=VideoStatus.uploaded.value,
|
||||
)
|
||||
session.add(video)
|
||||
await session.commit()
|
||||
await session.refresh(video)
|
||||
|
||||
return VideoOut.model_validate(video)
|
||||
|
||||
|
||||
@router.delete("/{video_id}")
|
||||
async def delete_video(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).where(Video.id == video_id))
|
||||
video = result.scalar_one_or_none()
|
||||
if not video:
|
||||
raise HTTPException(404, "Video not found")
|
||||
|
||||
if os.path.exists(video.filepath):
|
||||
os.remove(video.filepath)
|
||||
if video.output_path:
|
||||
output_full = os.path.join(settings.video_dir, video.output_path)
|
||||
if os.path.exists(output_full):
|
||||
os.remove(output_full)
|
||||
|
||||
await session.delete(video)
|
||||
await session.commit()
|
||||
return {"ok": True}
|
||||
|
||||
|
||||
@router.post("/{video_id}/process", response_model=ProcessResponse)
|
||||
async def start_processing(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).where(Video.id == video_id))
|
||||
video = result.scalar_one_or_none()
|
||||
if not video:
|
||||
raise HTTPException(404, "Video not found")
|
||||
|
||||
if video.status not in (VideoStatus.uploaded.value, VideoStatus.error.value):
|
||||
raise HTTPException(400, f"Cannot process video in status '{video.status}'")
|
||||
|
||||
existing_job = await session.execute(
|
||||
select(Job).where(
|
||||
Job.video_id == video_id,
|
||||
Job.job_type == "process",
|
||||
Job.status.in_([JobStatus.pending.value, JobStatus.running.value]),
|
||||
)
|
||||
)
|
||||
if existing_job.scalar_one_or_none():
|
||||
raise HTTPException(400, "A process job is already running for this video")
|
||||
|
||||
job = Job(video_id=video_id, job_type="process", status=JobStatus.pending.value)
|
||||
session.add(job)
|
||||
await session.commit()
|
||||
await session.refresh(job)
|
||||
|
||||
video.status = VideoStatus.queued.value
|
||||
await session.commit()
|
||||
|
||||
process_video.delay(video_id=video_id, job_id=job.id)
|
||||
|
||||
return ProcessResponse(job_id=job.id, message="Processing started")
|
||||
|
||||
|
||||
@router.get("/{video_id}/transcript", response_model=TranscriptOut)
|
||||
async def get_transcript(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(
|
||||
select(Segment).where(Segment.video_id == video_id).order_by(Segment.start_time)
|
||||
)
|
||||
segments = result.scalars().all()
|
||||
return TranscriptOut(
|
||||
video_id=video_id,
|
||||
segments=[SegmentOut.model_validate(s) for s in segments],
|
||||
)
|
||||
|
||||
|
||||
@router.put("/{video_id}/transcript/{segment_id}", response_model=SegmentOut)
|
||||
async def update_segment(
|
||||
video_id: int,
|
||||
segment_id: int,
|
||||
update: SegmentUpdate,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
):
|
||||
result = await session.execute(
|
||||
select(Segment).where(Segment.id == segment_id, Segment.video_id == video_id)
|
||||
)
|
||||
seg = result.scalar_one_or_none()
|
||||
if not seg:
|
||||
raise HTTPException(404, "Segment not found")
|
||||
|
||||
if update.english_text is not None:
|
||||
seg.english_text = update.english_text
|
||||
if update.danish_text is not None:
|
||||
seg.danish_text = update.danish_text
|
||||
if update.speaker_label is not None:
|
||||
seg.speaker_label = update.speaker_label
|
||||
|
||||
await session.commit()
|
||||
await session.refresh(seg)
|
||||
return SegmentOut.model_validate(seg)
|
||||
|
||||
|
||||
@router.get("/{video_id}/file")
|
||||
async def stream_video(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).where(Video.id == video_id))
|
||||
video = result.scalar_one_or_none()
|
||||
if not video:
|
||||
raise HTTPException(404, "Video not found")
|
||||
if not os.path.exists(video.filepath):
|
||||
raise HTTPException(404, "Video file not found on disk")
|
||||
return FileResponse(video.filepath, media_type="video/mp4", filename=video.filename)
|
||||
|
||||
|
||||
@router.get("/file/{filename}")
|
||||
async def stream_output(filename: str):
|
||||
full_path = os.path.join(settings.video_dir, filename)
|
||||
if not os.path.exists(full_path):
|
||||
raise HTTPException(404, "File not found")
|
||||
return FileResponse(full_path)
|
||||
|
||||
|
||||
@router.post("/{video_id}/generate", response_model=ProcessResponse)
|
||||
async def start_generation(video_id: int, session: AsyncSession = Depends(get_session)):
|
||||
result = await session.execute(select(Video).where(Video.id == video_id))
|
||||
video = result.scalar_one_or_none()
|
||||
if not video:
|
||||
raise HTTPException(404, "Video not found")
|
||||
|
||||
if video.status != VideoStatus.review.value:
|
||||
raise HTTPException(400, f"Cannot generate audio in status '{video.status}'")
|
||||
|
||||
existing_job = await session.execute(
|
||||
select(Job).where(
|
||||
Job.video_id == video_id,
|
||||
Job.job_type == "generate",
|
||||
Job.status.in_([JobStatus.pending.value, JobStatus.running.value]),
|
||||
)
|
||||
)
|
||||
if existing_job.scalar_one_or_none():
|
||||
raise HTTPException(400, "A generate job is already running for this video")
|
||||
|
||||
job = Job(video_id=video_id, job_type="generate", status=JobStatus.pending.value)
|
||||
session.add(job)
|
||||
await session.commit()
|
||||
await session.refresh(job)
|
||||
|
||||
video.status = VideoStatus.generating.value
|
||||
await session.commit()
|
||||
|
||||
generate_audio.delay(video_id=video_id, job_id=job.id)
|
||||
|
||||
return ProcessResponse(job_id=job.id, message="Audio generation started")
|
||||
@@ -0,0 +1,34 @@
|
||||
from pydantic_settings import BaseSettings
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
video_dir: str = "/data/videos"
|
||||
app_data_dir: str = "/data/app"
|
||||
database_url: str = "sqlite+aiosqlite:////data/app/bornetime.db"
|
||||
|
||||
redis_url: str = "redis://redis:6379/0"
|
||||
|
||||
stt_model_id: str = "bosonai/higgs-audio-v3-stt"
|
||||
stt_device: str = "cuda:0"
|
||||
stt_dtype: str = "bfloat16"
|
||||
|
||||
translate_model_path: str = "/models/translategemma-12b-q8_0.gguf"
|
||||
translate_n_gpu_layers: int = -1
|
||||
translate_device: str = "cuda:0"
|
||||
|
||||
tts_model_id: str = "bosonai/higgs-tts-3-4b"
|
||||
tts_device: str = "cuda:1"
|
||||
|
||||
diarization_model: str = "pyannote/speaker-diarization-3.1"
|
||||
diarization_device: str = "cuda:0"
|
||||
|
||||
hf_token: Optional[str] = None
|
||||
temp_dir: str = "/data/app/tmp"
|
||||
|
||||
class Config:
|
||||
env_file = ".env"
|
||||
env_file_encoding = "utf-8"
|
||||
|
||||
|
||||
settings = Settings()
|
||||
@@ -0,0 +1,43 @@
|
||||
import os
|
||||
import logging
|
||||
from contextlib import asynccontextmanager
|
||||
from fastapi import FastAPI
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
|
||||
from .models.database import init_db
|
||||
from .api.videos import router as videos_router
|
||||
from .api.jobs import router as jobs_router
|
||||
from .config import settings
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
os.makedirs(settings.video_dir, exist_ok=True)
|
||||
os.makedirs(settings.app_data_dir, exist_ok=True)
|
||||
os.makedirs(settings.temp_dir, exist_ok=True)
|
||||
await init_db()
|
||||
yield
|
||||
|
||||
|
||||
app = FastAPI(title="B\u00f8rnetime", version="1.0.0", lifespan=lifespan)
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=["*"],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
app.include_router(videos_router, prefix="/api/videos", tags=["videos"])
|
||||
app.include_router(jobs_router, prefix="/api/jobs", tags=["jobs"])
|
||||
|
||||
|
||||
@app.get("/api/health")
|
||||
async def health():
|
||||
return {"status": "ok"}
|
||||
@@ -0,0 +1,84 @@
|
||||
import logging
|
||||
import torch
|
||||
from typing import Optional
|
||||
from pyannote.audio import Pipeline
|
||||
from ..config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_diarization_pipeline: Optional[Pipeline] = None
|
||||
|
||||
|
||||
def get_diarization_pipeline() -> Pipeline:
|
||||
global _diarization_pipeline
|
||||
if _diarization_pipeline is None:
|
||||
import os
|
||||
|
||||
hf_token = settings.hf_token or os.environ.get("HUGGINGFACE_TOKEN")
|
||||
logger.info("Loading diarization model on %s", settings.diarization_device)
|
||||
_diarization_pipeline = Pipeline.from_pretrained(
|
||||
settings.diarization_model,
|
||||
use_auth_token=hf_token,
|
||||
)
|
||||
_diarization_pipeline.to(torch.device(settings.diarization_device))
|
||||
return _diarization_pipeline
|
||||
|
||||
|
||||
def diarize(audio_path: str) -> list[dict]:
|
||||
pipeline = get_diarization_pipeline()
|
||||
diarization = pipeline(audio_path)
|
||||
|
||||
segments = []
|
||||
for turn, _, speaker in diarization.itertracks(yield_label=True):
|
||||
segments.append(
|
||||
{
|
||||
"speaker_label": speaker,
|
||||
"start_time": turn.start,
|
||||
"end_time": turn.end,
|
||||
}
|
||||
)
|
||||
|
||||
segments.sort(key=lambda s: s["start_time"])
|
||||
segments = _merge_adjacent_same_speaker(segments)
|
||||
segments = _relabel_speakers(segments)
|
||||
|
||||
return segments
|
||||
|
||||
|
||||
def _merge_adjacent_same_speaker(
|
||||
segments: list[dict], max_gap: float = 0.5
|
||||
) -> list[dict]:
|
||||
if not segments:
|
||||
return []
|
||||
|
||||
merged = [segments[0]]
|
||||
for seg in segments[1:]:
|
||||
prev = merged[-1]
|
||||
if (
|
||||
prev["speaker_label"] == seg["speaker_label"]
|
||||
and (seg["start_time"] - prev["end_time"]) <= max_gap
|
||||
):
|
||||
prev["end_time"] = seg["end_time"]
|
||||
else:
|
||||
merged.append(seg)
|
||||
return merged
|
||||
|
||||
|
||||
def _relabel_speakers(segments: list[dict]) -> list[dict]:
|
||||
label_map = {}
|
||||
counter = 0
|
||||
for seg in segments:
|
||||
spk = seg["speaker_label"]
|
||||
if spk not in label_map:
|
||||
label = chr(65 + counter)
|
||||
label_map[spk] = f"Speaker {label}"
|
||||
counter += 1
|
||||
seg["speaker_label"] = label_map[spk]
|
||||
return segments
|
||||
|
||||
|
||||
def unload_model():
|
||||
global _diarization_pipeline
|
||||
_diarization_pipeline = None
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,110 @@
|
||||
import subprocess
|
||||
import os
|
||||
import numpy as np
|
||||
import soundfile as sf
|
||||
from typing import Optional
|
||||
from ..config import settings
|
||||
|
||||
|
||||
def mix_audio(
|
||||
video_path: str,
|
||||
segments: list[dict],
|
||||
output_path: str,
|
||||
):
|
||||
temp_dir = settings.temp_dir
|
||||
|
||||
original_audio_path = os.path.join(temp_dir, "original_audio.wav")
|
||||
_extract_audio(video_path, original_audio_path)
|
||||
|
||||
original_audio, orig_sr = sf.read(original_audio_path)
|
||||
if len(original_audio.shape) > 1:
|
||||
original_audio = original_audio.mean(axis=1)
|
||||
|
||||
target_sr = 24000
|
||||
if orig_sr != target_sr:
|
||||
import librosa
|
||||
|
||||
original_audio = librosa.resample(
|
||||
original_audio, orig_sr=orig_sr, target_sr=target_sr
|
||||
)
|
||||
|
||||
total_samples = len(original_audio)
|
||||
danish_track = np.zeros(total_samples, dtype=np.float32)
|
||||
|
||||
for seg in segments:
|
||||
danish_audio = seg.get("danish_audio")
|
||||
if danish_audio is None or len(danish_audio) == 0:
|
||||
continue
|
||||
|
||||
start_sample = int(seg["start_time"] * target_sr)
|
||||
end_sample = min(start_sample + len(danish_audio), total_samples)
|
||||
actual_len = end_sample - start_sample
|
||||
danish_track[start_sample:end_sample] += danish_audio[:actual_len]
|
||||
|
||||
volume_envelope = np.ones(total_samples, dtype=np.float32)
|
||||
for seg in segments:
|
||||
start_sample = int(seg["start_time"] * target_sr)
|
||||
end_sample = int(seg["end_time"] * target_sr)
|
||||
volume_envelope[start_sample:end_sample] = 0.23
|
||||
|
||||
orig_adjusted = original_audio * volume_envelope
|
||||
|
||||
mixed = orig_adjusted + danish_track
|
||||
peak = np.max(np.abs(mixed))
|
||||
if peak > 0.99:
|
||||
mixed = mixed / peak * 0.95
|
||||
|
||||
mixed_path = os.path.join(temp_dir, "mixed_audio.wav")
|
||||
sf.write(mixed_path, mixed, target_sr)
|
||||
|
||||
_replace_audio(video_path, mixed_path, output_path)
|
||||
|
||||
for f in [original_audio_path, mixed_path]:
|
||||
if os.path.exists(f):
|
||||
os.remove(f)
|
||||
|
||||
|
||||
def _extract_audio(video_path: str, output_path: str):
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vn",
|
||||
"-acodec",
|
||||
"pcm_s16le",
|
||||
"-ar",
|
||||
"24000",
|
||||
"-ac",
|
||||
"1",
|
||||
output_path,
|
||||
]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Audio extraction failed: {result.stderr}")
|
||||
|
||||
|
||||
def _replace_audio(video_path: str, audio_path: str, output_path: str):
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-i",
|
||||
audio_path,
|
||||
"-c:v",
|
||||
"copy",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"192k",
|
||||
"-map",
|
||||
"0:v:0",
|
||||
"-map",
|
||||
"1:a:0",
|
||||
"-shortest",
|
||||
output_path,
|
||||
]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Audio replacement failed: {result.stderr}")
|
||||
@@ -0,0 +1,141 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import soundfile as sf
|
||||
import librosa
|
||||
import re
|
||||
import logging
|
||||
from functools import partial
|
||||
from dataclasses import asdict
|
||||
from typing import Optional
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
AutoModel,
|
||||
AutoTokenizer,
|
||||
WhisperProcessor,
|
||||
)
|
||||
from boson_multimodal.data_collator.higgs_audio_collator import HiggsAudioSampleCollator
|
||||
from boson_multimodal.data_types import ChatMLSample, AudioContent, Message
|
||||
from boson_multimodal.dataset.chatml_dataset import (
|
||||
ChatMLDatasetSample,
|
||||
prepare_chatml_sample_qwen,
|
||||
)
|
||||
|
||||
from ..config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_model: Optional[AutoModel] = None
|
||||
_tokenizer: Optional[AutoTokenizer] = None
|
||||
_collator: Optional[HiggsAudioSampleCollator] = None
|
||||
_config: Optional[AutoConfig] = None
|
||||
|
||||
|
||||
def _load_model():
|
||||
global _model, _tokenizer, _collator, _config
|
||||
if _model is not None:
|
||||
return
|
||||
|
||||
logger.info("Loading Higgs STT on %s", settings.stt_device)
|
||||
dtype = torch.bfloat16 if settings.stt_dtype == "bfloat16" else torch.float16
|
||||
|
||||
_config = AutoConfig.from_pretrained(settings.stt_model_id, trust_remote_code=True)
|
||||
_model = AutoModel.from_pretrained(
|
||||
settings.stt_model_id,
|
||||
torch_dtype=dtype,
|
||||
trust_remote_code=True,
|
||||
attn_implementation="eager",
|
||||
device_map=settings.stt_device,
|
||||
)
|
||||
_model.eval()
|
||||
_tokenizer = AutoTokenizer.from_pretrained(settings.stt_model_id)
|
||||
_model.audio_out_bos_token_id = _tokenizer.convert_tokens_to_ids(
|
||||
"<|audio_out_bos|>"
|
||||
)
|
||||
_model.audio_eos_token_id = _tokenizer.convert_tokens_to_ids("<|audio_eos|>")
|
||||
|
||||
whisper_proc = WhisperProcessor.from_pretrained("openai/whisper-large-v3")
|
||||
_collator = HiggsAudioSampleCollator(
|
||||
whisper_processor=whisper_proc,
|
||||
audio_in_token_id=_config.audio_in_token_idx,
|
||||
audio_out_token_id=_config.audio_out_token_idx,
|
||||
audio_stream_bos_id=_config.audio_stream_bos_id,
|
||||
audio_stream_eos_id=_config.audio_stream_eos_id,
|
||||
encode_whisper_embed=_config.encode_whisper_embed,
|
||||
pad_token_id=_config.pad_token_id,
|
||||
return_audio_in_tokens=_config.encode_audio_in_tokens,
|
||||
use_delay_pattern=_config.use_delay_pattern,
|
||||
round_to=1,
|
||||
audio_num_codebooks=_config.audio_num_codebooks,
|
||||
chunk_size_seconds=getattr(_config, "chunk_size_seconds", 30),
|
||||
encoder_padding_method=getattr(_config, "encoder_padding_method", "max_length"),
|
||||
)
|
||||
|
||||
|
||||
def transcribe_segment(audio_path: str, start_time: float, end_time: float) -> str:
|
||||
_load_model()
|
||||
|
||||
audio_np, sr = sf.read(audio_path)
|
||||
if sr != 16000:
|
||||
audio_np = librosa.resample(audio_np, orig_sr=sr, target_sr=16000)
|
||||
|
||||
start_sample = int(start_time * 16000)
|
||||
end_sample = int(end_time * 16000)
|
||||
segment_audio = audio_np[start_sample:end_sample]
|
||||
|
||||
if len(segment_audio) == 0:
|
||||
return ""
|
||||
|
||||
prompt = "Transcribe the speech. Output only the spoken words in lowercase with no punctuation."
|
||||
messages = [
|
||||
Message(role="user", content=[prompt, AudioContent(audio_url="placeholder")])
|
||||
]
|
||||
chatml = ChatMLSample(messages=messages)
|
||||
prep_fn = partial(prepare_chatml_sample_qwen, enable_thinking=True)
|
||||
input_tokens, _, _, _ = prep_fn(chatml, _tokenizer, add_generation_prompt=True)
|
||||
|
||||
sample = ChatMLDatasetSample(
|
||||
input_ids=torch.LongTensor(input_tokens),
|
||||
label_ids=None,
|
||||
audio_ids_concat=None,
|
||||
audio_ids_start=None,
|
||||
audio_waveforms_concat=torch.tensor(segment_audio, dtype=torch.float32),
|
||||
audio_waveforms_start=torch.tensor([0]),
|
||||
audio_sample_rate=torch.tensor([16000]),
|
||||
audio_speaker_indices=torch.tensor([0]),
|
||||
)
|
||||
|
||||
batch = asdict(_collator([sample]))
|
||||
device = next(_model.parameters()).device
|
||||
batch = {
|
||||
k: v.to(device).contiguous() if isinstance(v, torch.Tensor) else v
|
||||
for k, v in batch.items()
|
||||
}
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = _model.generate(
|
||||
**batch,
|
||||
max_new_tokens=1024,
|
||||
use_cache=True,
|
||||
do_sample=False,
|
||||
stop_strings=["<|im_end|>", "<|endoftext|>"],
|
||||
tokenizer=_tokenizer,
|
||||
)
|
||||
|
||||
output_ids = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
full_text = _tokenizer.decode(output_ids[0], skip_special_tokens=False)
|
||||
|
||||
parts = full_text.split("assistant\n")
|
||||
hyp = parts[-1] if len(parts) > 1 else full_text
|
||||
hyp = re.sub(r"<think>.*?</think>", "", hyp, flags=re.DOTALL)
|
||||
hyp = re.sub(r"<\|.*?\|>", "", hyp).strip()
|
||||
|
||||
return hyp
|
||||
|
||||
|
||||
def unload_model():
|
||||
global _model, _tokenizer, _collator, _config
|
||||
_model = None
|
||||
_tokenizer = None
|
||||
_collator = None
|
||||
_config = None
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,60 @@
|
||||
import logging
|
||||
from ..config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_llm = None
|
||||
|
||||
|
||||
def _load_model():
|
||||
global _llm
|
||||
if _llm is not None:
|
||||
return
|
||||
|
||||
from llama_cpp import Llama
|
||||
|
||||
logger.info("Loading TranslateGemma 12B from %s", settings.translate_model_path)
|
||||
_llm = Llama(
|
||||
model_path=settings.translate_model_path,
|
||||
n_gpu_layers=settings.translate_n_gpu_layers,
|
||||
n_ctx=2048,
|
||||
verbose=False,
|
||||
)
|
||||
|
||||
|
||||
def translate(text: str) -> str:
|
||||
_load_model()
|
||||
|
||||
if not text.strip():
|
||||
return ""
|
||||
|
||||
lines = [
|
||||
"Translate the following English text to Danish.",
|
||||
"Preserve the original meaning, tone, and style.",
|
||||
"Output only the Danish translation, nothing else.",
|
||||
"",
|
||||
f"English: {text}",
|
||||
"",
|
||||
"Danish:",
|
||||
]
|
||||
prompt = "\n".join(lines)
|
||||
|
||||
stop_tokens = ["\n\n", "English:", "User:"]
|
||||
output = _llm(
|
||||
prompt,
|
||||
max_tokens=1024,
|
||||
temperature=0.1,
|
||||
stop=stop_tokens,
|
||||
echo=False,
|
||||
)
|
||||
|
||||
result = output["choices"][0]["text"].strip()
|
||||
return result
|
||||
|
||||
|
||||
def unload_model():
|
||||
global _llm
|
||||
_llm = None
|
||||
import torch
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,256 @@
|
||||
import io
|
||||
import os
|
||||
import logging
|
||||
import numpy as np
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from ..config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_model = None
|
||||
_tokenizer = None
|
||||
_use_http = True
|
||||
|
||||
|
||||
def _load_http_mode():
|
||||
global _use_http
|
||||
import requests
|
||||
|
||||
try:
|
||||
resp = requests.get(
|
||||
settings.tts_api_url.replace("/v1/audio/speech", "/health"), timeout=5
|
||||
)
|
||||
if resp.ok:
|
||||
_use_http = True
|
||||
logger.info("TTS sidecar available at %s", settings.tts_api_url)
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
_use_http = False
|
||||
return False
|
||||
|
||||
|
||||
def _load_model():
|
||||
global _model, _tokenizer
|
||||
if _model is not None:
|
||||
return
|
||||
|
||||
logger.info("Loading Higgs TTS 3 on %s", settings.tts_device)
|
||||
from transformers import AutoModel, AutoTokenizer
|
||||
|
||||
_model = AutoModel.from_pretrained(
|
||||
settings.tts_model_id,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=True,
|
||||
device_map=settings.tts_device,
|
||||
)
|
||||
_model.eval()
|
||||
|
||||
_tokenizer = AutoTokenizer.from_pretrained(settings.tts_model_id)
|
||||
logger.info("Higgs TTS 3 loaded successfully")
|
||||
|
||||
|
||||
def synthesize_segment(
|
||||
text: str,
|
||||
reference_audio_path: Optional[str] = None,
|
||||
reference_text: Optional[str] = None,
|
||||
speaker_label: Optional[str] = None,
|
||||
) -> bytes:
|
||||
if not text.strip():
|
||||
return b""
|
||||
|
||||
input_text = text
|
||||
if speaker_label:
|
||||
input_text = f"<|emotion:contentment|><|prosody:expressive_low|>{text}"
|
||||
|
||||
if _use_http:
|
||||
try:
|
||||
return _synthesize_http(input_text, reference_audio_path, reference_text)
|
||||
except Exception as e:
|
||||
logger.warning("HTTP TTS failed, falling back to direct: %s", e)
|
||||
|
||||
return _synthesize_direct(input_text, reference_audio_path, reference_text)
|
||||
|
||||
|
||||
def _synthesize_http(
|
||||
text: str,
|
||||
reference_audio_path: Optional[str] = None,
|
||||
reference_text: Optional[str] = None,
|
||||
) -> bytes:
|
||||
import requests
|
||||
|
||||
payload = {
|
||||
"input": text,
|
||||
"model": settings.tts_model_id,
|
||||
"temperature": 0.8,
|
||||
"top_k": 50,
|
||||
"max_new_tokens": 2048,
|
||||
}
|
||||
|
||||
if reference_audio_path and reference_text:
|
||||
import base64
|
||||
|
||||
with open(reference_audio_path, "rb") as f:
|
||||
audio_b64 = base64.b64encode(f.read()).decode()
|
||||
payload["references"] = [
|
||||
{
|
||||
"audio": audio_b64,
|
||||
"text": reference_text,
|
||||
}
|
||||
]
|
||||
|
||||
resp = requests.post(settings.tts_api_url, json=payload, timeout=120)
|
||||
resp.raise_for_status()
|
||||
return resp.content
|
||||
|
||||
|
||||
def _synthesize_direct(
|
||||
text: str,
|
||||
reference_audio_path: Optional[str] = None,
|
||||
reference_text: Optional[str] = None,
|
||||
) -> bytes:
|
||||
_load_model()
|
||||
|
||||
from boson_multimodal.data_types import ChatMLSample, Message
|
||||
from boson_multimodal.dataset.chatml_dataset import (
|
||||
ChatMLDatasetSample,
|
||||
prepare_chatml_sample_qwen,
|
||||
)
|
||||
from boson_multimodal.data_collator.higgs_audio_collator import (
|
||||
HiggsAudioSampleCollator,
|
||||
)
|
||||
from functools import partial
|
||||
from dataclasses import asdict
|
||||
|
||||
prompt = f"Generate speech for the following text. Output only the spoken words."
|
||||
messages = [Message(role="user", content=[prompt])]
|
||||
chatml = ChatMLSample(messages=messages)
|
||||
prep_fn = partial(prepare_chatml_sample_qwen, enable_thinking=False)
|
||||
input_tokens, _, _, _ = prep_fn(chatml, _tokenizer, add_generation_prompt=True)
|
||||
|
||||
ref_audio = None
|
||||
if reference_audio_path:
|
||||
import soundfile as sf
|
||||
|
||||
ref_audio, ref_sr = sf.read(reference_audio_path)
|
||||
if ref_sr != 16000:
|
||||
import librosa
|
||||
|
||||
ref_audio = librosa.resample(ref_audio, orig_sr=ref_sr, target_sr=16000)
|
||||
ref_audio = torch.tensor(ref_audio, dtype=torch.float32)
|
||||
|
||||
sample = ChatMLDatasetSample(
|
||||
input_ids=torch.LongTensor(input_tokens),
|
||||
label_ids=None,
|
||||
audio_ids_concat=None,
|
||||
audio_ids_start=None,
|
||||
audio_waveforms_concat=ref_audio if ref_audio is not None else torch.zeros(1),
|
||||
audio_waveforms_start=torch.tensor([0]),
|
||||
audio_sample_rate=torch.tensor([16000]),
|
||||
audio_speaker_indices=torch.tensor([0]),
|
||||
)
|
||||
|
||||
collator = HiggsAudioSampleCollator(
|
||||
whisper_processor=None,
|
||||
audio_in_token_id=getattr(_model.config, "audio_in_token_idx", 0),
|
||||
audio_out_token_id=getattr(_model.config, "audio_out_token_idx", 1),
|
||||
audio_stream_bos_id=getattr(_model.config, "audio_stream_bos_id", 2),
|
||||
audio_stream_eos_id=getattr(_model.config, "audio_stream_eos_id", 3),
|
||||
encode_whisper_embed=getattr(_model.config, "encode_whisper_embed", False),
|
||||
pad_token_id=_tokenizer.pad_token_id or 0,
|
||||
return_audio_in_tokens=getattr(_model.config, "encode_audio_in_tokens", False),
|
||||
use_delay_pattern=getattr(_model.config, "use_delay_pattern", True),
|
||||
round_to=1,
|
||||
audio_num_codebooks=getattr(_model.config, "audio_num_codebooks", 8),
|
||||
chunk_size_seconds=getattr(_model.config, "chunk_size_seconds", 30),
|
||||
encoder_padding_method=getattr(
|
||||
_model.config, "encoder_padding_method", "max_length"
|
||||
),
|
||||
)
|
||||
|
||||
batch = asdict(collator([sample]))
|
||||
device = next(_model.parameters()).device
|
||||
batch = {
|
||||
k: v.to(device).contiguous() if isinstance(v, torch.Tensor) else v
|
||||
for k, v in batch.items()
|
||||
}
|
||||
|
||||
with torch.inference_mode():
|
||||
outputs = _model.generate(
|
||||
**batch,
|
||||
max_new_tokens=2048,
|
||||
temperature=0.8,
|
||||
top_k=50,
|
||||
do_sample=True,
|
||||
use_cache=True,
|
||||
)
|
||||
|
||||
output_ids = outputs[0] if isinstance(outputs, tuple) else outputs
|
||||
|
||||
if hasattr(_tokenizer, "decode_audio"):
|
||||
audio_data = _tokenizer.decode_audio(output_ids)
|
||||
if isinstance(audio_data, bytes):
|
||||
return audio_data
|
||||
return audio_data.tobytes()
|
||||
|
||||
if hasattr(_model, "decode_audio"):
|
||||
audio_data = _model.decode_audio(output_ids)
|
||||
if isinstance(audio_data, bytes):
|
||||
return audio_data
|
||||
if isinstance(audio_data, torch.Tensor):
|
||||
audio_np = audio_data.cpu().float().numpy()
|
||||
return _numpy_to_wav_bytes(audio_np, 24000)
|
||||
|
||||
if isinstance(output_ids, torch.Tensor):
|
||||
audio_np = output_ids.cpu().float().numpy()
|
||||
if audio_np.ndim > 1:
|
||||
audio_np = audio_np.flatten()
|
||||
return _numpy_to_wav_bytes(audio_np, 24000)
|
||||
|
||||
return b""
|
||||
|
||||
|
||||
def _numpy_to_wav_bytes(audio: np.ndarray, sample_rate: int = 24000) -> bytes:
|
||||
import wave
|
||||
import struct
|
||||
|
||||
audio = np.clip(audio, -1.0, 1.0)
|
||||
audio_int16 = (audio * 32767).astype(np.int16)
|
||||
|
||||
buf = io.BytesIO()
|
||||
with wave.open(buf, "wb") as wf:
|
||||
wf.setnchannels(1)
|
||||
wf.setsampwidth(2)
|
||||
wf.setframerate(sample_rate)
|
||||
wf.writeframes(audio_int16.tobytes())
|
||||
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def _bytes_to_numpy(
|
||||
audio_bytes: bytes, target_sr: int = 24000
|
||||
) -> tuple[np.ndarray, int]:
|
||||
import soundfile as sf
|
||||
|
||||
with io.BytesIO(audio_bytes) as buf:
|
||||
data, sr = sf.read(buf)
|
||||
|
||||
if len(data.shape) > 1:
|
||||
data = data.mean(axis=1)
|
||||
|
||||
if sr != target_sr:
|
||||
import librosa
|
||||
|
||||
data = librosa.resample(data, orig_sr=sr, target_sr=target_sr)
|
||||
|
||||
return data.astype(np.float32), target_sr
|
||||
|
||||
|
||||
def unload_model():
|
||||
global _model, _tokenizer
|
||||
_model = None
|
||||
_tokenizer = None
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,47 @@
|
||||
import subprocess
|
||||
import os
|
||||
from typing import Optional
|
||||
from ..config import settings
|
||||
|
||||
|
||||
def get_video_duration(path: str) -> Optional[float]:
|
||||
cmd = [
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
path,
|
||||
]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return float(result.stdout.strip())
|
||||
return None
|
||||
|
||||
|
||||
def extract_audio(video_path: str, output_path: str, sample_rate: int = 16000) -> str:
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-vn",
|
||||
"-acodec",
|
||||
"pcm_s16le",
|
||||
"-ar",
|
||||
str(sample_rate),
|
||||
"-ac",
|
||||
"1",
|
||||
output_path,
|
||||
]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"Audio extraction failed: {result.stderr}")
|
||||
return output_path
|
||||
|
||||
|
||||
def build_output_path(video_path: str) -> str:
|
||||
base, ext = os.path.splitext(os.path.basename(video_path))
|
||||
return f"{base}_en_da_audio{ext}"
|
||||
@@ -0,0 +1,19 @@
|
||||
from celery import Celery
|
||||
from ..config import settings
|
||||
|
||||
celery_app = Celery(
|
||||
"bornetime",
|
||||
broker=settings.redis_url,
|
||||
backend=settings.redis_url,
|
||||
)
|
||||
|
||||
celery_app.conf.update(
|
||||
task_serializer="json",
|
||||
accept_content=["json"],
|
||||
result_serializer="json",
|
||||
timezone="UTC",
|
||||
enable_utc=True,
|
||||
task_track_started=True,
|
||||
task_acks_late=True,
|
||||
worker_prefetch_multiplier=1,
|
||||
)
|
||||
@@ -0,0 +1,189 @@
|
||||
import os
|
||||
import shutil
|
||||
import logging
|
||||
import numpy as np
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .celery_app import celery_app
|
||||
from ..config import settings
|
||||
from ..models.tables import Video, Segment, Job, VideoStatus, JobStatus
|
||||
from ..services.tts import synthesize_segment, unload_model as unload_tts
|
||||
from ..services.mixer import mix_audio
|
||||
from ..services.video_utils import build_output_path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
engine = create_engine(settings.database_url.replace("+aiosqlite", ""))
|
||||
|
||||
|
||||
def _get_sync_session():
|
||||
return Session(engine)
|
||||
|
||||
|
||||
@celery_app.task(bind=True, name="generate_audio")
|
||||
def generate_audio(self, video_id: int, job_id: int = None):
|
||||
db = _get_sync_session()
|
||||
temp_files = []
|
||||
audio_dir = None
|
||||
try:
|
||||
video = db.query(Video).filter(Video.id == video_id).first()
|
||||
if not video:
|
||||
raise ValueError(f"Video {video_id} not found")
|
||||
|
||||
video.status = VideoStatus.generating.value
|
||||
db.commit()
|
||||
|
||||
if job_id:
|
||||
job = db.query(Job).filter(Job.id == job_id).first()
|
||||
else:
|
||||
job = Job(
|
||||
video_id=video_id, job_type="generate", status=JobStatus.running.value
|
||||
)
|
||||
db.add(job)
|
||||
db.commit()
|
||||
|
||||
job.celery_task_id = self.request.id
|
||||
job.status = JobStatus.running.value
|
||||
db.commit()
|
||||
|
||||
segments = (
|
||||
db.query(Segment)
|
||||
.filter(Segment.video_id == video_id)
|
||||
.order_by(Segment.start_time)
|
||||
.all()
|
||||
)
|
||||
|
||||
total = len(segments)
|
||||
audio_dir = os.path.join(settings.temp_dir, f"tts_{video_id}")
|
||||
os.makedirs(audio_dir, exist_ok=True)
|
||||
|
||||
speaker_refs = {}
|
||||
for seg in segments:
|
||||
if seg.speaker_label and seg.speaker_label not in speaker_refs:
|
||||
seg_audio_path = os.path.join(
|
||||
settings.temp_dir,
|
||||
f"ref_{video_id}_{seg.speaker_label.replace(' ', '_')}.wav",
|
||||
)
|
||||
temp_files.append(seg_audio_path)
|
||||
_extract_segment_audio(
|
||||
video.filepath, seg.start_time, seg.end_time, seg_audio_path
|
||||
)
|
||||
speaker_refs[seg.speaker_label] = {
|
||||
"audio_path": seg_audio_path,
|
||||
"text": seg.english_text,
|
||||
}
|
||||
|
||||
logger.info("Starting TTS for %d segments", total)
|
||||
segment_data = []
|
||||
for i, seg in enumerate(segments):
|
||||
ref = speaker_refs.get(seg.speaker_label, {})
|
||||
audio_bytes = synthesize_segment(
|
||||
text=seg.danish_text or seg.english_text or "",
|
||||
reference_audio_path=ref.get("audio_path"),
|
||||
reference_text=ref.get("text"),
|
||||
speaker_label=seg.speaker_label,
|
||||
)
|
||||
|
||||
seg_audio_path = os.path.join(audio_dir, f"seg_{seg.id}.wav")
|
||||
with open(seg_audio_path, "wb") as f:
|
||||
f.write(audio_bytes)
|
||||
temp_files.append(seg_audio_path)
|
||||
|
||||
import soundfile as sf
|
||||
|
||||
danish_audio, sr = sf.read(seg_audio_path)
|
||||
if len(danish_audio.shape) > 1:
|
||||
danish_audio = danish_audio.mean(axis=1)
|
||||
if sr != 24000:
|
||||
import librosa
|
||||
|
||||
danish_audio = librosa.resample(
|
||||
danish_audio, orig_sr=sr, target_sr=24000
|
||||
)
|
||||
|
||||
segment_data.append(
|
||||
{
|
||||
"start_time": seg.start_time,
|
||||
"end_time": seg.end_time,
|
||||
"danish_audio": danish_audio.astype(np.float32),
|
||||
}
|
||||
)
|
||||
|
||||
pct = (i + 1) / total * 0.9
|
||||
self.update_state(state="PROGRESS", meta={"progress": pct})
|
||||
job.progress = pct
|
||||
db.commit()
|
||||
|
||||
logger.info("Mixing audio for video %d", video_id)
|
||||
output_path = build_output_path(video.filepath)
|
||||
mix_audio(video.filepath, segment_data, output_path)
|
||||
|
||||
video.output_path = output_path
|
||||
video.status = VideoStatus.done.value
|
||||
job.status = JobStatus.done.value
|
||||
job.progress = 1.0
|
||||
db.commit()
|
||||
|
||||
unload_tts()
|
||||
logger.info("Generation complete for video %d", video_id)
|
||||
return {"video_id": video_id, "output_path": output_path}
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Generation failed for video %d: %s", video_id, e, exc_info=True)
|
||||
video = db.query(Video).filter(Video.id == video_id).first()
|
||||
if video:
|
||||
video.status = VideoStatus.error.value
|
||||
video.error_message = str(e)
|
||||
if job_id:
|
||||
job = db.query(Job).filter(Job.id == job_id).first()
|
||||
else:
|
||||
job = (
|
||||
db.query(Job)
|
||||
.filter(Job.video_id == video_id, Job.job_type == "generate")
|
||||
.first()
|
||||
)
|
||||
if job:
|
||||
job.status = JobStatus.failed.value
|
||||
job.error_message = str(e)
|
||||
db.commit()
|
||||
raise
|
||||
|
||||
finally:
|
||||
for f in temp_files:
|
||||
if os.path.exists(f):
|
||||
try:
|
||||
os.remove(f)
|
||||
except OSError:
|
||||
pass
|
||||
if audio_dir and os.path.exists(audio_dir):
|
||||
try:
|
||||
shutil.rmtree(audio_dir)
|
||||
except OSError:
|
||||
pass
|
||||
db.close()
|
||||
|
||||
|
||||
def _extract_segment_audio(video_path: str, start: float, end: float, output_path: str):
|
||||
import subprocess
|
||||
|
||||
duration = end - start
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
video_path,
|
||||
"-ss",
|
||||
str(start),
|
||||
"-t",
|
||||
str(duration),
|
||||
"-vn",
|
||||
"-acodec",
|
||||
"pcm_s16le",
|
||||
"-ar",
|
||||
"16000",
|
||||
"-ac",
|
||||
"1",
|
||||
output_path,
|
||||
]
|
||||
subprocess.run(cmd, capture_output=True, check=True)
|
||||
@@ -0,0 +1,144 @@
|
||||
import os
|
||||
import logging
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .celery_app import celery_app
|
||||
from ..config import settings
|
||||
from ..models.tables import Video, Segment, Job, VideoStatus, JobStatus
|
||||
from ..services.diarization import diarize, unload_model as unload_diarization
|
||||
from ..services.stt import transcribe_segment, unload_model as unload_stt
|
||||
from ..services.translator import translate, unload_model as unload_translate
|
||||
from ..services.video_utils import extract_audio
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
engine = create_engine(settings.database_url.replace("+aiosqlite", ""))
|
||||
|
||||
|
||||
def _get_sync_session():
|
||||
return Session(engine)
|
||||
|
||||
|
||||
@celery_app.task(bind=True, name="process_video")
|
||||
def process_video(self, video_id: int, job_id: int = None):
|
||||
db = _get_sync_session()
|
||||
temp_files = []
|
||||
try:
|
||||
video = db.query(Video).filter(Video.id == video_id).first()
|
||||
if not video:
|
||||
raise ValueError(f"Video {video_id} not found")
|
||||
|
||||
video.status = VideoStatus.processing.value
|
||||
db.commit()
|
||||
|
||||
if job_id:
|
||||
job = db.query(Job).filter(Job.id == job_id).first()
|
||||
else:
|
||||
job = Job(
|
||||
video_id=video_id, job_type="process", status=JobStatus.running.value
|
||||
)
|
||||
db.add(job)
|
||||
db.commit()
|
||||
|
||||
self.update_state(state="PROGRESS", meta={"progress": 0.0})
|
||||
job.celery_task_id = self.request.id
|
||||
job.status = JobStatus.running.value
|
||||
db.commit()
|
||||
|
||||
audio_path = os.path.join(settings.temp_dir, f"audio_{video_id}.wav")
|
||||
temp_files.append(audio_path)
|
||||
extract_audio(video.filepath, audio_path)
|
||||
logger.info("Extracted audio for video %d", video_id)
|
||||
|
||||
self.update_state(state="PROGRESS", meta={"progress": 0.10})
|
||||
job.progress = 0.10
|
||||
db.commit()
|
||||
|
||||
segments = diarize(audio_path)
|
||||
if not segments:
|
||||
raise RuntimeError("No speech segments detected in video")
|
||||
logger.info("Diarized %d segments for video %d", len(segments), video_id)
|
||||
|
||||
unload_diarization()
|
||||
|
||||
self.update_state(state="PROGRESS", meta={"progress": 0.20})
|
||||
job.progress = 0.20
|
||||
db.commit()
|
||||
|
||||
total = len(segments)
|
||||
|
||||
logger.info("Starting STT for %d segments", total)
|
||||
for i, seg in enumerate(segments):
|
||||
en_text = transcribe_segment(audio_path, seg["start_time"], seg["end_time"])
|
||||
seg["english_text"] = en_text
|
||||
pct = 0.20 + (i + 1) / total * 0.35
|
||||
self.update_state(state="PROGRESS", meta={"progress": pct})
|
||||
job.progress = pct
|
||||
db.commit()
|
||||
|
||||
unload_stt()
|
||||
logger.info("STT complete, unloaded model")
|
||||
|
||||
logger.info("Starting translation for %d segments", total)
|
||||
for i, seg in enumerate(segments):
|
||||
en_text = seg.get("english_text", "")
|
||||
da_text = translate(en_text) if en_text else ""
|
||||
seg["danish_text"] = da_text
|
||||
|
||||
pct = 0.55 + (i + 1) / total * 0.35
|
||||
self.update_state(state="PROGRESS", meta={"progress": pct})
|
||||
job.progress = pct
|
||||
db.commit()
|
||||
|
||||
unload_translate()
|
||||
logger.info("Translation complete, unloaded model")
|
||||
|
||||
for seg in segments:
|
||||
db_seg = Segment(
|
||||
video_id=video_id,
|
||||
speaker_label=seg["speaker_label"],
|
||||
start_time=seg["start_time"],
|
||||
end_time=seg["end_time"],
|
||||
english_text=seg.get("english_text", ""),
|
||||
danish_text=seg.get("danish_text", ""),
|
||||
status="done",
|
||||
)
|
||||
db.add(db_seg)
|
||||
|
||||
video.status = VideoStatus.review.value
|
||||
job.status = JobStatus.done.value
|
||||
job.progress = 1.0
|
||||
db.commit()
|
||||
|
||||
logger.info("Processing complete for video %d", video_id)
|
||||
return {"video_id": video_id, "status": "review"}
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Processing failed for video %d: %s", video_id, e, exc_info=True)
|
||||
video = db.query(Video).filter(Video.id == video_id).first()
|
||||
if video:
|
||||
video.status = VideoStatus.error.value
|
||||
video.error_message = str(e)
|
||||
if job_id:
|
||||
job = db.query(Job).filter(Job.id == job_id).first()
|
||||
else:
|
||||
job = (
|
||||
db.query(Job)
|
||||
.filter(Job.video_id == video_id, Job.job_type == "process")
|
||||
.first()
|
||||
)
|
||||
if job:
|
||||
job.status = JobStatus.failed.value
|
||||
job.error_message = str(e)
|
||||
db.commit()
|
||||
raise
|
||||
|
||||
finally:
|
||||
for f in temp_files:
|
||||
if os.path.exists(f):
|
||||
try:
|
||||
os.remove(f)
|
||||
except OSError:
|
||||
pass
|
||||
db.close()
|
||||
@@ -0,0 +1,19 @@
|
||||
fastapi>=0.115.0
|
||||
uvicorn[standard]>=0.30.0
|
||||
sqlalchemy[asyncio]>=2.0.0
|
||||
aiosqlite>=0.20.0
|
||||
pydantic>=2.0.0
|
||||
pydantic-settings>=2.0.0
|
||||
celery>=5.4.0
|
||||
redis>=5.0.0
|
||||
python-multipart>=0.0.12
|
||||
ffmpeg-python>=0.2.0
|
||||
torch>=2.4.0
|
||||
transformers>=4.51.0
|
||||
boson_multimodal
|
||||
soundfile>=0.12.0
|
||||
librosa>=0.10.0
|
||||
pyannote.audio>=3.0.0
|
||||
llama-cpp-python>=0.3.0
|
||||
numpy>=1.24.0
|
||||
aiofiles>=24.0.0
|
||||
Reference in New Issue
Block a user