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:
2026-07-10 22:42:23 +02:00
commit 6851b26923
43 changed files with 3390 additions and 0 deletions
View File
View File
+54
View File
@@ -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")
+218
View File
@@ -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")
+34
View File
@@ -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()
+43
View File
@@ -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"}
View File
+84
View File
@@ -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()
+110
View File
@@ -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}")
+141
View File
@@ -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()
+60
View File
@@ -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()
+256
View File
@@ -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()
+47
View File
@@ -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}"
View File
+19
View File
@@ -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,
)
+189
View File
@@ -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)
+144
View File
@@ -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()
+19
View File
@@ -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