eigyouonsei / app.py
ryota
GPUが混んでいても止まらないようにし、出力をExcelだけにする
7fb75a2
Raw
History Blame Contribute Delete
23.1 kB
"""商談文字起こし(Gradio / ZeroGPU)
HuggingFace Spaces の Gradio SDK で動かす。PRO に付いている ZeroGPU を使うと、
文字起こしと話者分離が GPU で走る。GPU が無い環境(自分のPCなど)でも
そのまま CPU で動く。
起動:
python app.py
"""
from __future__ import annotations
import os
import shutil
import tempfile
import time
import traceback
from pathlib import Path
import gradio as gr
import pipeline as pl
# ---------------------------------------------------------------- ZeroGPU
# Spaces では spaces パッケージが入っている。自分のPCには無いので、
# 無ければ「何もしない飾り」に差し替えて同じコードが動くようにする。
try:
import spaces # type: ignore
HAS_SPACES = True
except Exception: # 自分のPCには無い。Spaces でも環境次第で読めないことがある
spaces = None
HAS_SPACES = False
ON_SPACES = os.environ.get("SPACE_ID") is not None
ON_ZERO_GPU = HAS_SPACES and os.environ.get("SPACES_ZERO_GPU") is not None
GPU_DEVICE = "cuda" if ON_ZERO_GPU else ""
# GPU を確保する時間の候補。長く要求するほど空きが見つかりにくいので、
# 録音の長さに合わせていちばん短いもので頼む。
GPU_TIERS = (60, 120, 300)
# 順番待ちで取れなかったときに何回まで粘るか
GPU_RETRIES = int(os.environ.get("GPU_RETRIES", "2"))
def on_gpu(seconds: int):
"""ZeroGPU のときだけ GPU を割り当てる。それ以外は素通し。"""
def deco(fn):
if not ON_ZERO_GPU:
return fn
return spaces.GPU(duration=seconds)(fn)
return deco
def carry_errors(fn):
"""GPU側で起きた失敗を、文字にして持ち帰る。
ZeroGPU は別プロセスで動かした結果を送り返すが、例外はうまく復元できず
「'RuntimeError'」のような中身の無い形になって、原因が分からなくなる。
そこで成否と本文を組にして返し、呼び出し側で組み立て直す。
"""
def wrapped(*args, **kwargs):
try:
return True, fn(*args, **kwargs)
except Exception as exc:
traceback.print_exc()
return False, f"{type(exc).__name__}: {exc}"
return wrapped
def unwrap(result):
ok, payload = result
if not ok:
raise RuntimeError(payload)
return payload
def _transcribe(wav, model_size, prompt, fast, device):
return pl.transcribe_file(Path(wav), model_size, prompt, fast, device=device)
def _diarize(wav, hf_token, num_speakers, device):
return pl.diarize(Path(wav), hf_token, num_speakers, device=device)
# 確保時間ごとに用意しておく。ZeroGPU の指定は関数を作るときに決まるため。
TRANSCRIBE_GPU = {
seconds: on_gpu(seconds)(carry_errors(
lambda wav, model_size, prompt, fast: _transcribe(wav, model_size, prompt, fast, "cuda")
))
for seconds in GPU_TIERS
}
DIARIZE_GPU = {
seconds: on_gpu(seconds)(carry_errors(
lambda wav, hf_token, num_speakers: _diarize(wav, hf_token, num_speakers, "cuda")
))
for seconds in GPU_TIERS
}
def gpu_tier(wav, factor: float) -> int:
"""この録音に必要そうな確保時間。モデル読込ぶんを足して見積もる。"""
try:
length = pl.probe_duration(Path(wav))
except Exception:
length = 0.0
needed = 30 + length * factor
for seconds in GPU_TIERS:
if needed <= seconds:
return seconds
return GPU_TIERS[-1]
def is_busy(message: str) -> bool:
"""GPUの順番待ちで弾かれたか(故障ではなく、混んでいるだけ)。"""
text = message.lower()
return "no gpu" in text or "gpu was available" in text or "quota" in text
def run_on_gpu(table, wav, args, factor, cpu_call, note):
"""GPUで実行し、混んでいて取れなければCPUに落とす。
失敗させるより、遅くても最後まで終わったほうが役に立つ。
"""
seconds = gpu_tier(wav, factor)
last = ""
for attempt in range(GPU_RETRIES + 1):
ok, payload = table[seconds](wav, *args)
if ok:
return payload
last = str(payload)
if not is_busy(last):
raise RuntimeError(last) # 混雑ではない本当の失敗
print(f"[info] GPUの空き待ち({attempt + 1}回目): {last}")
time.sleep(5)
print(f"[info] GPUが取れないのでCPUで処理します: {last}")
note("GPUが混んでいたため、一部はCPUで処理しました(時間がかかります)")
return cpu_call()
def make_workers(note):
"""pipeline に渡す、文字起こしと話者分離の実体を作る。"""
def transcriber(wav, model_size, prompt, fast):
if not ON_ZERO_GPU:
return _transcribe(wav, model_size, prompt, fast, GPU_DEVICE)
return run_on_gpu(
TRANSCRIBE_GPU, wav, (model_size, prompt, fast), 0.15,
lambda: _transcribe(wav, model_size, prompt, fast, "cpu"), note,
)
def diarizer(wav, hf_token, num_speakers):
if not ON_ZERO_GPU:
return _diarize(wav, hf_token, num_speakers, GPU_DEVICE)
return run_on_gpu(
DIARIZE_GPU, wav, (hf_token, num_speakers), 0.20,
lambda: _diarize(wav, hf_token, num_speakers, "cpu"), note,
)
return transcriber, diarizer
# ---------------------------------------------------------------- 置き場所
WORK_ROOT = Path(tempfile.gettempdir()) / "spinthoughts"
WORK_ROOT.mkdir(parents=True, exist_ok=True)
# 商談の音声を必要以上に置いておかないため、古いものは消す
JOB_TTL_SECONDS = 6 * 3600
def sweep_old_jobs() -> None:
now = time.time()
for path in WORK_ROOT.iterdir():
try:
if path.is_dir() and now - path.stat().st_mtime > JOB_TTL_SECONDS:
shutil.rmtree(path, ignore_errors=True)
except OSError:
continue
# ---------------------------------------------------------------- 画面に出す形
STATUS_MARK = {"transcribed": "", "skipped": "除外", "error": "失敗"}
def esc(text) -> str:
return (
str(text)
.replace("&", "&amp;")
.replace("<", "&lt;")
.replace(">", "&gt;")
.replace('"', "&quot;")
)
def audio_url(path: Path) -> str:
return "/gradio_api/file=" + Path(path).as_posix()
def render(results) -> str:
"""録音ごとのカード。時刻を押すとその位置から音声が鳴る。"""
if not results:
return ""
# 時刻が押せることは見ただけでは分からないので、最初に書いておく
guide = ""
if any(r.playback for r in results):
guide = (
'<p class="guide">発話の左にある<b>時刻</b>を押すと、'
'その場面から録音が流れます。聞き取りにくい行だけ確かめられます。</p>'
)
cards = []
for i, entry in enumerate(results, 1):
audio_id = "audio-%d" % i
tag = STATUS_MARK.get(entry.status, "")
meta = " · ".join(x for x in [
pl.hhmmss(entry.duration),
entry.recorded_at,
("%d発話" % len(entry.utterances)) if entry.status == "transcribed" else "",
entry.reason or "",
] if x)
player = ""
if entry.playback:
player = (
'<audio id="%s" class="player" controls preload="none" src="%s"></audio>'
% (audio_id, audio_url(entry.playback))
)
talk = pl.speaking_time(entry)
talk_html = ""
if len(talk) > 1:
parts = [
"<span><b>%s</b> %d分%02d秒・%d回</span>"
% (esc(name), int(secs // 60), int(secs % 60), count)
for name, secs, count in talk
]
talk_html = '<div class="talk">%s</div>' % "".join(parts)
palette = {}
lines = []
for u in entry.utterances:
if u.speaker not in palette:
palette[u.speaker] = len(palette) % 5
flag = ""
if u.needs_review:
flag = '<span class="flag" title="%s">要確認</span>' % esc(u.review_note)
who = '<span class="who">%s</span>' % esc(u.speaker) if u.speaker else ""
lines.append(
'<div class="line">'
'<button class="seek" data-audio="%s" data-t="%.2f">%s</button>'
'<span class="spine sp%d"></span>'
'<span class="said">%s<span class="what">%s%s</span></span>'
"</div>"
% (audio_id, u.start, pl.hhmmss(u.start), palette[u.speaker],
who, esc(u.text), flag)
)
if lines:
body = '<div class="score">%s</div>' % "".join(lines)
else:
body = '<p class="empty">%s</p>' % esc(
entry.reason or "発話を検出できませんでした。"
)
cards.append(
'<details class="file"%s>'
'<summary><span class="no">%03d</span>'
'<span class="file-name">%s</span>%s'
'<span class="file-meta">%s</span></summary>'
'<div class="file-body">%s%s%s</div>'
"</details>"
% (
" open" if entry.status == "transcribed" else "",
i,
esc(os.path.basename(entry.original_name)),
('<span class="tag">%s</span>' % tag) if tag else "",
esc(meta),
player, talk_html, body,
)
)
return '<div class="results">%s%s</div>' % (guide, "".join(cards))
def summarize(results) -> str:
transcribed = sum(1 for r in results if r.status == "transcribed")
skipped = sum(1 for r in results if r.status == "skipped")
failed = sum(1 for r in results if r.status == "error")
flagged = sum(1 for r in results for u in r.utterances if u.needs_review)
return (
"**書き起こし %d** / 長さ不足で除外 %d / 失敗 %d / 要確認 %d"
% (transcribed, skipped, failed, flagged)
)
# ---------------------------------------------------------------- 実行
def process(files, min_seconds, model_size, num_speakers, diarization, fast, prompt,
hf_token, progress=gr.Progress()):
if not files:
raise gr.Error("音声ファイルかZIPを選んでください。")
token = (hf_token or "").strip() or os.environ.get("HF_TOKEN", "")
if diarization and not token:
raise gr.Error(
"話者分離にはHuggingFaceのトークンが必要です。"
"トークンを入れるか、話者分離をオフにしてください。"
)
sweep_old_jobs()
workdir = Path(tempfile.mkdtemp(prefix="job-", dir=WORK_ROOT))
# Gradio が置く一時ファイルは名前が変わることがあるので、元の名前で置き直す
upload_dir = workdir / "受け取り"
upload_dir.mkdir(parents=True, exist_ok=True)
sources = []
for item in files:
src = Path(item if isinstance(item, str) else item.name)
target = upload_dir / pl.safe_filename(src.name)
counter = 1
while target.exists():
target = upload_dir / ("%s__%d%s" % (target.stem, counter, target.suffix))
counter += 1
shutil.copyfile(src, target)
sources.append(target)
def notify(message, done, total):
progress((done, total) if total else 0, desc=message)
# GPUが取れずCPUに落ちたときなど、利用者に伝えるべきことを溜める
notes: list[str] = []
def note(message):
if message not in notes:
notes.append(message)
transcriber, diarizer = make_workers(note)
try:
results = pl.run(
sources=sources,
workdir=workdir,
min_seconds=max(0.0, float(min_seconds or 0)),
model_size=model_size,
num_speakers=int(num_speakers) if num_speakers and int(num_speakers) > 0 else None,
diarization_enabled=bool(diarization),
fast=bool(fast),
prompt=prompt or "",
hf_token=token,
progress=notify,
transcriber=transcriber,
diarizer=diarizer,
)
except Exception as exc:
traceback.print_exc()
raise gr.Error(str(exc))
progress(0, desc="ファイルを書き出しています")
outputs = pl.build_outputs(results, workdir / "出力", workdir.name, xlsx_only=True)
# 元の音声と作業用ファイルは消す。聞き返す用(再生用)だけ残す。
shutil.rmtree(workdir / "audio", ignore_errors=True)
shutil.rmtree(workdir / "wav", ignore_errors=True)
shutil.rmtree(upload_dir, ignore_errors=True)
downloads = [str(outputs["xlsx"])] if outputs.get("xlsx") else []
lines = [summarize(results)]
lines += [f"※ {n}" for n in notes]
return " \n".join(lines), render(results), downloads
# ---------------------------------------------------------------- 画面
HEAD = """
<script>
/* 時刻を押したら、その位置から音声を鳴らす。
結果は後から差し込まれるので、document 側でまとめて受ける。 */
document.addEventListener("click", function (e) {
var btn = e.target.closest("button.seek");
if (!btn) return;
e.preventDefault();
var audio = document.getElementById(btn.dataset.audio);
if (!audio) return;
audio.currentTime = parseFloat(btn.dataset.t || "0");
audio.play().catch(function () {});
});
</script>
"""
CSS = """
/* 明るい画面と暗い画面のどちらでも読めるように、色は変数で持つ */
/* 設定欄。説明文の行数が違っても、入力欄の位置を揃える。
横並びを壊さないよう、各項目の入れ物だけを対象にする。 */
#settings .form { align-items: stretch; }
#settings .form > .block { display: flex; flex-direction: column; }
#settings .form > .block > :last-child { margin-top: auto; }
.results {
--card: #ffffff;
--line: #dbe3ee;
--ink: #16202c;
--soft: #5b6b7c;
--faint: #8695a6;
--accent: #1d5fd0;
--accent-soft: #eaf1fd;
--warn: #9a5b0e;
--warn-soft: #fdf0dc;
--sp0: #1d5fd0; --sp1: #c2410c; --sp2: #6d28d9; --sp3: #0f766e; --sp4: #b91c1c;
font-feature-settings: "palt";
}
/* 暗い画面。Gradio の切り替え(.dark)と、OSの設定の両方に反応させる */
.dark .results,
.dark.results {
--card: #131c27;
--line: #2b3746;
--ink: #e8eef6;
--soft: #a3b1c2;
--faint: #8494a5;
--accent: #7fb2ff;
--accent-soft: #1b2b42;
--warn: #e0a44a;
--warn-soft: #35290f;
--sp0: #7fb2ff; --sp1: #f2955c; --sp2: #b39dfa; --sp3: #4fc3ae; --sp4: #f47b7b;
}
@media (prefers-color-scheme: dark) {
.results {
--card: #131c27;
--line: #2b3746;
--ink: #e8eef6;
--soft: #a3b1c2;
--faint: #8494a5;
--accent: #7fb2ff;
--accent-soft: #1b2b42;
--warn: #e0a44a;
--warn-soft: #35290f;
--sp0: #7fb2ff; --sp1: #f2955c; --sp2: #b39dfa; --sp3: #4fc3ae; --sp4: #f47b7b;
}
}
.results { color: var(--ink); }
.results .guide {
font-size: 13px; color: var(--soft);
background: var(--accent-soft); border-radius: 10px;
padding: 10px 14px; margin: 0 0 12px;
}
.results .guide b { color: var(--accent); }
.results .file {
background: var(--card);
border: 1px solid var(--line);
border-radius: 12px;
margin-bottom: 12px;
overflow: hidden;
}
.results .file[open] { border-color: var(--accent); }
.results summary {
padding: 14px 16px;
cursor: pointer;
display: flex;
gap: 12px;
align-items: baseline;
flex-wrap: wrap;
list-style: none;
}
.results summary::-webkit-details-marker { display: none; }
.results summary:hover { background: var(--accent-soft); }
.results .no {
font-family: ui-monospace, SFMono-Regular, Menlo, monospace;
font-size: 11px;
color: #fff;
background: var(--accent);
padding: 2px 8px;
border-radius: 20px;
flex: none;
}
.results .file-name { font-weight: 700; font-size: 15px; }
.results .file-meta { font-size: 12px; color: var(--faint); margin-left: auto; }
.results .tag {
font-size: 11px; padding: 2px 9px; border-radius: 20px;
background: var(--accent-soft); color: var(--accent); font-weight: 600;
}
.results .file-body { padding: 4px 16px 16px; }
.results .player {
width: 100%; height: 38px; margin-bottom: 12px;
border-radius: 8px;
}
.results .talk {
display: flex; gap: 8px; flex-wrap: wrap; margin-bottom: 12px;
}
.results .talk span {
font-size: 12px; color: var(--soft);
background: var(--accent-soft); border-radius: 20px; padding: 4px 12px;
}
.results .talk b { color: var(--ink); font-weight: 600; }
.results .score { border-top: 1px solid var(--line); }
.results .line {
display: flex; gap: 12px; align-items: flex-start;
padding: 12px 0;
border-bottom: 1px solid var(--line);
}
.results .line:last-child { border-bottom: none; }
/* 時刻は押すと頭出しできる。下線ではなく、押せる形で示す */
.results button.seek {
font-family: ui-monospace, SFMono-Regular, Menlo, monospace;
font-size: 11.5px;
color: var(--accent);
background: var(--accent-soft);
border: none;
border-radius: 6px;
padding: 4px 8px;
cursor: pointer;
text-decoration: none;
flex: none;
transition: background .12s, color .12s;
}
.results button.seek:hover { background: var(--accent); color: #fff; }
.results .spine { width: 3px; border-radius: 3px; flex: none; align-self: stretch; background: var(--faint); }
.results .sp0 { background: var(--sp0); } .results .sp1 { background: var(--sp1); }
.results .sp2 { background: var(--sp2); } .results .sp3 { background: var(--sp3); }
.results .sp4 { background: var(--sp4); }
.results .said { flex: 1; min-width: 0; }
.results .who {
display: block;
font-size: 11px; font-weight: 700; letter-spacing: .04em;
color: var(--soft);
margin-bottom: 3px;
}
.results .what { font-size: 15px; line-height: 1.9; word-break: break-word; color: var(--ink); }
.results .flag {
display: inline-block;
font-size: 10.5px; color: var(--warn); background: var(--warn-soft);
padding: 1px 8px; border-radius: 20px; margin-left: 8px; vertical-align: 2px;
}
.results .empty { color: var(--faint); font-size: 13px; padding: 8px 0; }
@media (max-width: 620px) {
.results .file-meta { margin-left: 0; width: 100%; }
.results .what { font-size: 14.5px; }
}
"""
AUDIO_TYPES = [
".zip", ".wav", ".mp3", ".m4a", ".mp4", ".aac", ".flac", ".ogg", ".opus",
".wma", ".amr", ".3gp", ".mov", ".aif", ".aiff",
]
DEFAULTS = {
"min_seconds": 60,
"model_size": "large-v3-turbo",
"num_speakers": 2,
"diarization": False,
"fast": True,
}
def environment_note() -> str:
if ON_ZERO_GPU:
return "**ZeroGPU で動作中。** 処理のたびにGPUが割り当てられます。"
if pl.gpu_label():
return "**%s で動作中。**" % pl.gpu_label()
return (
"**CPU処理(%dスレッド)。** 録音1時間あたり30〜60分かかります。"
% pl.cpu_threads()
)
with gr.Blocks(title="商談文字起こし") as demo:
gr.Markdown("# 商談文字起こし")
gr.Markdown(
"音声ファイルやZIPを入れると、一定の長さ以上の録音だけを文字起こししてCSVにします。 \n"
+ environment_note()
)
saved = gr.BrowserState(DEFAULTS, storage_key="spinthoughts.settings.v3")
files = gr.File(
label="音声ファイル または ZIP(複数可)",
file_count="multiple",
file_types=AUDIO_TYPES,
)
with gr.Row(elem_id="settings"):
min_seconds = gr.Number(
DEFAULTS["min_seconds"], label="この長さ以上だけ処理する(秒)",
info="これより短い録音は文字起こしせず、一覧にだけ残します",
)
model_size = gr.Dropdown(
[("標準(turbo・速くて実用精度)", "large-v3-turbo"),
("速度優先(small・誤変換が増えます)", "small"),
("精度優先(large-v3・遅い)", "large-v3")],
value=DEFAULTS["model_size"], label="精度",
info="迷ったら標準のままで構いません",
)
num_speakers = gr.Number(
DEFAULTS["num_speakers"], label="話者の人数",
info="話者を分けるときだけ使います。0で自動判定",
)
with gr.Row():
diarization = gr.Checkbox(DEFAULTS["diarization"], label="話者を分けて記録する")
fast = gr.Checkbox(DEFAULTS["fast"], label="速さを優先する")
prompt = gr.Textbox(
pl.DEFAULT_PROMPT, label="よく出る言葉", lines=3,
info="業務でよく使う語を書いておくと、固有名詞や専門用語の精度が上がります",
)
hf_token = gr.Textbox(
"", label="HuggingFaceトークン", type="password",
visible=not os.environ.get("HF_TOKEN"),
info="話者分離に必要です。Space の Secret に HF_TOKEN があれば表示されません",
)
run = gr.Button("文字起こしを始める", variant="primary")
summary = gr.Markdown()
downloads = gr.Files(label="出力(録音ごと.xlsx / 1録音=1シート)")
results = gr.HTML()
run.click(
process,
inputs=[files, min_seconds, model_size, num_speakers, diarization, fast,
prompt, hf_token],
outputs=[summary, results, downloads],
api_name='transcribe',
)
# --- 設定を覚える --------------------------------------------------------
settings = [min_seconds, model_size, num_speakers, diarization, fast]
def restore(store):
store = store or {}
return [store.get(key, value) for key, value in DEFAULTS.items()]
def remember(*values):
return dict(zip(DEFAULTS.keys(), values))
demo.load(restore, inputs=saved, outputs=settings)
for component in settings:
component.change(remember, inputs=settings, outputs=saved)
def listen_on() -> str:
"""待ち受け先。
自分のPCでは 127.0.0.1(外から入れない)。Spaces では外側の入口から
コンテナ内へ届かないと起動失敗になるので 0.0.0.0 にする。
"""
if os.environ.get("HOST"):
return os.environ["HOST"]
return "0.0.0.0" if ON_SPACES else "127.0.0.1"
if __name__ == "__main__":
password = os.environ.get("APP_PASSWORD")
demo.launch(
server_name=listen_on(),
server_port=int(os.environ.get("PORT", "7860")),
auth=(os.environ.get("APP_USER", "spin"), password) if password else None,
allowed_paths=[str(WORK_ROOT)],
head=HEAD,
css=CSS,
theme=gr.themes.Soft(primary_hue="blue", neutral_hue="slate"),
)