File size: 5,016 Bytes
ace6209
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""OpenAI-compatible transcription server for GigaAM-He.

    pip install fastapi uvicorn soundfile torch torchaudio numpy
    pip install git+https://github.com/salute-developers/GigaAM.git
    python server.py                       # :6002

    curl -s http://localhost:6002/v1/audio/transcriptions \
         -F file=@clip.wav | jq -r .text

The shape of the response matches OpenAI's /v1/audio/transcriptions, so any
client that speaks to Whisper's API speaks to this: point its base URL here.
"""
import io
import logging
import os
import sys
import time

import numpy as np
import soundfile as sf
import torch
import uvicorn
from fastapi import FastAPI, File, Form, HTTPException, UploadFile

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from longform import SAMPLE_RATE, segment_audio  # noqa: E402

CHECKPOINT = os.getenv("HE_CHECKPOINT", "gigaam-he-twostage.ckpt")
PORT = int(os.getenv("HE_PORT", "6002"))
DEVICE = os.getenv("HE_DEVICE", "cuda" if torch.cuda.is_available() else "cpu")
MAX_LEN = float(os.getenv("HE_MAX_LEN", "18.0"))   # the model is trained on <=20s
BATCH = int(os.getenv("HE_BATCH", "8"))

logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
app = FastAPI(title="GigaAM-He")
_model = None


def _load():
    global _model
    import gigaam
    logging.info(f"loading {CHECKPOINT} on {DEVICE}")
    _model = gigaam.load_model(CHECKPOINT, device=DEVICE)
    logging.info("model ready")


def _warmup():
    """Touch every length bucket once so no caller pays algorithm selection.

    Inputs are bucketed to whole seconds, so there are only ~20 possible
    shapes. Compiling them all at startup removes a ~190 ms first-touch spike
    from the serving path -- otherwise the first request of each new length
    pays it, which is exactly the wrong place for it.
    """
    t0 = time.time()
    for secs in range(1, int(MAX_LEN) + 2):
        buf = io.BytesIO()
        sf.write(buf, np.zeros(int(secs * SAMPLE_RATE), dtype="float32"),
                 SAMPLE_RATE, format="WAV")
        _transcribe(buf.getvalue())
    logging.info(f"warmup done ({int(MAX_LEN) + 1} buckets, {time.time() - t0:.1f}s)")


def _transcribe(audio_bytes: bytes) -> str:
    import torchaudio.functional as taF

    x, sr = sf.read(io.BytesIO(audio_bytes), dtype="float32")
    if x.ndim > 1:
        x = x.mean(axis=1)
    if sr != SAMPLE_RATE:
        # torchaudio's resampler, never np.interp: linear interpolation is not
        # an anti-aliasing filter and folds everything above 8 kHz back into
        # the band.
        x = taF.resample(torch.from_numpy(x), sr, SAMPLE_RATE).numpy()
    if len(x) < int(0.1 * SAMPLE_RATE):
        return ""

    spans = segment_audio(x, max_len_s=MAX_LEN)
    chunks = [torch.from_numpy(x[int(s * SAMPLE_RATE):int(e * SAMPLE_RATE)]) for s, e in spans]
    chunks = [c for c in chunks if len(c) >= int(0.1 * SAMPLE_RATE)]
    if not chunks:
        return ""

    out = []
    with torch.inference_mode():
        for i in range(0, len(chunks), BATCH):
            b = chunks[i:i + BATCH]
            lens = torch.tensor([len(c) for c in b], dtype=torch.long)
            # Round the padded width up to a whole second. Every distinct input
            # length is a distinct tensor shape and cuDNN re-runs algorithm
            # selection per shape -- measured, that costs 6x (p50 181 ms across
            # varied lengths vs 28 ms when one shape repeats). Bucketing
            # collapses thousands of widths to ~20. True lengths still go to the
            # model, so the padding changes nothing about the output.
            width = int(lens.max())
            width = -(-width // SAMPLE_RATE) * SAMPLE_RATE
            pad = torch.zeros(len(b), width, dtype=torch.float32)
            for j, c in enumerate(b):
                pad[j, :len(c)] = c
            enc, enc_len = _model.forward(
                pad.to(_model._device).to(_model._dtype), lens.to(_model._device))
            for text, _ in _model._decode(enc, enc_len, lens.to(_model._device), False):
                if text.strip():
                    out.append(text.strip())
    return " ".join(out)


@app.post("/v1/audio/transcriptions")
async def transcriptions(file: UploadFile = File(...),
                         model: str = Form(None),
                         language: str = Form(None)):
    """`model` and `language` are accepted and ignored: this server hosts one
    Hebrew model, and rejecting the fields would break OpenAI clients that
    always send them."""
    try:
        return {"text": _transcribe(await file.read())}
    except Exception as e:  # noqa: BLE001
        logging.exception("transcription failed")
        raise HTTPException(status_code=400, detail=str(e))


@app.get("/health")
def health():
    return {"status": "ok", "model": os.path.basename(CHECKPOINT), "device": DEVICE}


if __name__ == "__main__":
    _load()
    _warmup()
    uvicorn.run(app, host=os.getenv("HE_HOST", "0.0.0.0"), port=PORT)