From a18f8432250bd72d04dcdeb0600a3119b572657c Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 24 Aug 2026 14:27:14 +0000 Subject: [PATCH 1/3] Prime Whisper with each video's own words, and fix what it still mishears MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Reading the 83 scripts already in Drive, the errors are not random: they are the words a general model has no reason to expect. 费半 comes out as 飞班 or 肺瓣, 对冲基金 as 对中基金 (10+ scripts), 资产负债表 as 自然负债表 (6), 每股收益 as 美股收益, SNDK as S&DK, ISRG as ISIG, and Chinese sentences are punctuated with ASCII commas. Three layers, all of them cheap: - prompt_from_metadata (on) puts the video's title and the first line of its description in front of the model. The title already names the day's tickers. - vocabulary adds the terms a channel says every episode and rewrites what still comes out wrong. zh-finance ships for a Mandarin US-market channel; ASCII rewrites only match whole words, so CTS => CDS leaves CTSX alone. - polish (on) gives Chinese sentences fullwidth punctuation while leaving the commas in 1,250 alone, keeps one copy of a looped phrase, and drops subtitle boilerplate. whisper_condition_on_previous_text now defaults to false. Whisper's own default feeds each clip the previous clip's text, which is what makes it loop; batched decoding ignores the flag anyway, so off also keeps the batched path and the out-of-memory retry sounding the same. convert_to_simplified runs OpenCC over the text when the zh extra is installed, and warns rather than failing when it is not. `ytscript polish PATH...` re-applies the clean-up to scripts already written, so adding a term does not mean re-transcribing a backlog. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01GARmHk71niEh1TPunhH3Z8 --- README.md | 84 +++++++++- pyproject.toml | 3 + src/ytscript/cli.py | 85 +++++++++- src/ytscript/config.py | 58 ++++++- src/ytscript/data/zh-finance.txt | 157 +++++++++++++++++ src/ytscript/pipeline.py | 29 +++- src/ytscript/polish.py | 137 +++++++++++++++ src/ytscript/transcribers/__init__.py | 2 + src/ytscript/transcribers/base.py | 8 +- src/ytscript/transcribers/faster_whisper.py | 22 ++- src/ytscript/transcribers/openai_api.py | 7 +- src/ytscript/vocabulary.py | 177 ++++++++++++++++++++ tests/fakes.py | 24 ++- tests/test_cli.py | 31 ++++ tests/test_pipeline.py | 42 +++++ tests/test_polish.py | 84 ++++++++++ tests/test_vocabulary.py | 99 +++++++++++ 17 files changed, 1021 insertions(+), 28 deletions(-) create mode 100644 src/ytscript/data/zh-finance.txt create mode 100644 src/ytscript/polish.py create mode 100644 src/ytscript/vocabulary.py create mode 100644 tests/test_polish.py create mode 100644 tests/test_vocabulary.py diff --git a/README.md b/README.md index 3309701..9b292b8 100644 --- a/README.md +++ b/README.md @@ -34,7 +34,8 @@ machine. The `openai` extra needs `OPENAI_API_KEY` in the environment and charge minute of audio, but needs no local model. They are declared as conflicting extras — two transcription stacks, nothing needs both — so sync one at a time. The `drive` extra is independent of both and adds the Google client libraries for [saving the scripts to -Google Drive](#google-drive). +Google Drive](#google-drive), as is the `zh` extra, which adds [OpenCC] for rewriting +traditional characters as simplified. **On an NVIDIA GPU, install the CUDA libraries too.** faster-whisper runs on [CTranslate2], which needs cuBLAS and cuDNN 9 and does not bundle them. In a checkout: @@ -77,6 +78,8 @@ ytscript drive-auth # optional: sign in to Google Drive once ytscript run --dry-run # what would be transcribed ytscript run # first run: the latest 30 videos ytscript run # later runs: only what is new + +ytscript polish scripts # re-clean scripts already written ``` Scripts land in `output_dir` as `2024-05-01_Video-title_VIDEOID.txt`, and every finished @@ -91,6 +94,7 @@ Useful flags on `run`: | `--channel @handle` | Override the configured channel | | `--language zh` | Override the main language (`auto` detects it) | | `--batch-size N` | Clips decoded at once; `1` turns batching off | +| `--vocabulary NAME` | Terms and rewrites for the channel: a built-in name or a file | | `--limit N` | Check the newest N videos, whatever the state file says | | `--backfill` | Check `initial_backfill` videos again (default 30) | | `--format txt,md,json` | Write more than one rendering | @@ -106,6 +110,11 @@ Useful flags on `run`: The cookie and members-only flags work on `list` too. +`ytscript polish` runs the same clean-up a run does over scripts that already exist — +useful after adding a term to the vocabulary, or on a backlog transcribed before it had +one. It takes files or directories (`.txt` and `.md`), rewrites them in place, and has +`--vocabulary NAME`, `--simplified` and `--dry-run`. + ## Configuration Settings are read from `ytscript.toml` (searched for in the working directory and its @@ -125,11 +134,17 @@ whisper_device = "cuda" # "cpu", "cuda", ... whisper_compute_type = "float16" # "int8" on CPU, "float16" on GPU whisper_initial_prompt = "以下是普通话的句子。" # seeds simplified characters whisper_batch_size = 4 # clips decoded at once; 1 turns batching off +whisper_condition_on_previous_text = false # false stops the model looping a phrase + +prompt_from_metadata = true # prime each video with its own title and description +vocabulary = "zh-finance" # terms and rewrites; a built-in name or a file path output_dir = "scripts" output_formats = ["txt"] # any of txt, md, json timestamps = false paragraph_gap = 2.0 # silence in seconds that starts a new paragraph +polish = true # tidy punctuation, looped phrases and known bad spellings +convert_to_simplified = false # traditional -> simplified; needs --extra zh state_file = ".ytscript-state.json" keep_audio = false @@ -224,14 +239,73 @@ whisper_initial_prompt = "以下是普通话的句子。" # simplified whisper_initial_prompt = "以下是普通話的句子。" # traditional ``` -It is a nudge, not a guarantee — a few characters can still come out the other way. If -the output has to be uniform, run a converter such as [OpenCC] over the scripts -afterwards. The prompt also matters if you change `language`: a Chinese seed on English -audio makes the transcription worse, so change or remove it along with the language. +It is a nudge, not a guarantee — a few characters can still come out the other way. For +uniform output, `uv sync --extra zh` and set `convert_to_simplified = true`, which runs +[OpenCC] over the text before it is written; without the extra installed the setting +warns once and leaves the characters alone. The prompt also matters if you change +`language`: a Chinese seed on English audio makes the transcription worse, so change or +remove it along with the language. Leaving `whisper_initial_prompt` out of the config +entirely gets the stock sentence for `language`; `whisper_initial_prompt = ""` gets no +seed at all. Output is written with no spaces between Chinese segments, the way the language is written; a Latin word inside a sentence still keeps the spaces on either side of it. +### Getting the words right + +Whisper decodes a word it is expecting far more readily than one it is not, and the +words a niche channel leans on — tickers, index nicknames, the host's own name — are +exactly the ones a general model has no reason to expect. Left alone it guesses at the +sounds: 费半 (a Mandarin nickname for the Philadelphia semiconductor index) comes out as +飞班 or 肺瓣, 对冲基金 as 对中基金, 资产负债表 as 自然负债表, SNDK as `S&DK`. + +Three settings work on this, and they stack: + +- **`prompt_from_metadata`** (on by default) puts the video's own title and the first + line of its description in front of the model as the transcript so far. A title like + `半导体、ASML、TSM、NFLX、ISRG、MU、SNDK` names most of the day's tickers before a word + of audio is decoded, and it costs nothing — the metadata is already downloaded. +- **`vocabulary`** adds the terms the channel says every episode, and rewrites what the + model gets wrong anyway. `"zh-finance"` ships with ytscript for a Mandarin US-market + channel; the file is small, commented and meant to be copied: + + ```bash + cp "$(uv run python -c 'import ytscript.vocabulary as v; print(v.DATA_DIR)')/zh-finance.txt" mine.txt + ``` + + ``` + # a term is seeded into the prompt + 杰克逊霍尔 + # a rewrite is seeded and applied to the output + 对中基金 => 对冲基金 + CTS => CDS + ``` + + Rewrites are literal, not regular expressions. An ASCII left-hand side only matches as + a whole word, so `CTS` never fires inside `CTSX`; a Chinese one matches anywhere, since + the language has no word boundaries to key on. Point `vocabulary` at your copy and add + to it as you read the scripts — the mistakes a model makes are specific to the voice. + +- **`polish`** (on by default) cleans the text after recognition: Chinese sentences get + the fullwidth punctuation they are written with (`大家好,` → `大家好,`), while the + commas in `1,250` and the colon in a URL are left alone; a phrase the model looped on + is kept once instead of a dozen times; and the vocabulary's rewrites are applied. + +The prompt is capped at 200 characters, because Whisper conditions on about 224 tokens +and silently drops the rest. The seed sentence and the video's own subject go in first, +then terms — the ones named in the title before the rest, since those are the words the +episode actually says. + +`whisper_condition_on_previous_text` is the other accuracy setting, and ytscript +defaults it to `false` where Whisper's own default is `true`. Feeding each clip the text +of the one before it is what makes the model repeat a phrase for a minute once the audio +goes quiet. Batched decoding has no previous clip to condition on and behaves as if the +flag were off regardless, so `false` also keeps batched and sequential runs — including +the one-clip-at-a-time retry after an out-of-memory error — sounding the same. + +None of this touches the audio, so a script can be re-cleaned without re-transcribing: +add the term, then `ytscript polish scripts`. + ### Choosing a model `large-v3` at `float16` is the default because it is the best Chinese accuracy an 8 GB diff --git a/pyproject.toml b/pyproject.toml index 9892ef4..c1d953b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,6 +19,9 @@ dependencies = [ local = ["faster-whisper>=1.1.0"] # Hosted speech-to-text via the OpenAI audio transcription endpoint. openai = ["openai>=1.30.0"] +# Traditional -> simplified conversion, for `convert_to_simplified`. Pure Python, +# and independent of which speech-to-text backend is in use. +zh = ["opencc-python-reimplemented>=0.1.7"] # Copy the finished scripts into Google Drive. Independent of the two backends. # PySocks is not optional in practice: without it httplib2, which the Google client # is built on, quietly ignores every proxy setting instead of using one. diff --git a/src/ytscript/cli.py b/src/ytscript/cli.py index dd438d4..8389d1a 100644 --- a/src/ytscript/cli.py +++ b/src/ytscript/cli.py @@ -11,6 +11,8 @@ from .drive import DriveError, DriveUploader from .models import RunReport from .pipeline import Pipeline +from .polish import polish_text +from .vocabulary import VocabularyError, load_vocabulary from .youtube import YouTubeError log = logging.getLogger("ytscript") @@ -76,6 +78,10 @@ def add_common(target: argparse.ArgumentParser) -> None: type=int, help="clips decoded at once; 1 turns batching off", ) + run.add_argument( + "--vocabulary", + help="terms and rewrites for this channel: a built-in name or a file path", + ) run.add_argument("--output-dir", dest="output_dir", type=Path) run.add_argument( "--format", @@ -112,6 +118,31 @@ def add_common(target: argparse.ArgumentParser) -> None: help="show what would be transcribed without downloading anything", ) + polish = sub.add_parser( + "polish", + help="re-apply the vocabulary and punctuation clean-up to scripts already written", + ) + polish.add_argument( + "paths", + nargs="+", + type=Path, + metavar="PATH", + help="script files, or directories to take the .txt and .md files from", + ) + polish.add_argument("--vocabulary", help="a built-in name or a file path") + polish.add_argument( + "--simplified", + dest="convert_to_simplified", + action="store_true", + default=None, + help="also rewrite traditional characters as simplified", + ) + polish.add_argument( + "--dry-run", + action="store_true", + help="list the files that would change without writing them", + ) + listing = sub.add_parser("list", help="show the newest videos on the channel") add_common(listing) listing.add_argument("--limit", type=int, default=10) @@ -134,6 +165,7 @@ def add_common(target: argparse.ArgumentParser) -> None: "backend", "whisper_model", "whisper_batch_size", + "vocabulary", "output_dir", "output_formats", "timestamps", @@ -211,6 +243,56 @@ def cmd_list(args: argparse.Namespace) -> int: return 0 +SCRIPT_SUFFIXES = (".txt", ".md") + + +def _script_paths(paths: list[Path]) -> list[Path]: + """Expand the arguments into the script files to rewrite.""" + found: list[Path] = [] + for path in paths: + if path.is_dir(): + found.extend( + sorted(child for child in path.iterdir() if child.suffix in SCRIPT_SUFFIXES) + ) + elif path.is_file(): + found.append(path) + else: + raise ConfigError(f"no such file or directory: {path}") + return found + + +def cmd_polish(args: argparse.Namespace) -> int: + # No channel is needed to rewrite files that already exist. + config = load_config(path=args.config) + name = args.vocabulary if args.vocabulary is not None else config.vocabulary + vocabulary = load_vocabulary(name) + simplified = ( + config.convert_to_simplified + if args.convert_to_simplified is None + else args.convert_to_simplified + ) + paths = _script_paths(args.paths) + if not paths: + print("no .txt or .md scripts found") + return 0 + + changed = [] + for path in paths: + original = path.read_text(encoding="utf-8") + updated = polish_text(original, vocabulary, simplified=simplified) + if updated == original: + continue + changed.append(path) + if not args.dry_run: + path.write_text(updated, encoding="utf-8") + + verb = "would rewrite" if args.dry_run else "rewrote" + print(f"{verb} {len(changed)} of {len(paths)} file(s)") + for path in changed: + print(f" {path}") + return 0 + + def cmd_drive_auth(args: argparse.Namespace) -> int: # No channel is needed to authorise, so this skips the usual validation. config = load_config(path=args.config) @@ -251,12 +333,13 @@ def main(argv: list[str] | None = None) -> int: handlers = { "run": cmd_run, "list": cmd_list, + "polish": cmd_polish, "init": cmd_init, "drive-auth": cmd_drive_auth, } try: return handlers[args.command](args) - except (ConfigError, YouTubeError, DriveError) as exc: + except (ConfigError, VocabularyError, YouTubeError, DriveError) as exc: print(f"error: {exc}", file=sys.stderr) return 1 except KeyboardInterrupt: # pragma: no cover diff --git a/src/ytscript/config.py b/src/ytscript/config.py index 3220045..643930e 100644 --- a/src/ytscript/config.py +++ b/src/ytscript/config.py @@ -8,6 +8,8 @@ from pathlib import Path from typing import Any +from .vocabulary import VocabularyError, load_vocabulary + DEFAULT_CONFIG_NAMES = ("ytscript.toml", ".ytscript.toml") ENV_PREFIX = "YTSCRIPT_" @@ -45,7 +47,22 @@ class Config: """Defaults target an 8 GB NVIDIA card; see the README for the CPU-only settings.""" whisper_initial_prompt: str | None = None - """Seed text that steers spelling and register — for Chinese, simplified vs traditional.""" + """Seed text that steers spelling and register — for Chinese, simplified vs traditional. + + ``None`` uses the stock sentence for ``language``; set it to ``""`` for no seed.""" + + whisper_condition_on_previous_text: bool = False + """Feed each clip the previous clip's text. Whisper's own default is on, and it is + what makes the model repeat a phrase for a minute once it starts. Batched decoding + turns it off regardless, so leaving it off also keeps the two paths in step.""" + + prompt_from_metadata: bool = True + """Prime each video with its own title and description, so the day's tickers and + names are words the model is already expecting.""" + + vocabulary: str | None = None + """Terms the channel says every episode, and rewrites for what the model gets wrong. + A built-in name (``"zh-finance"``) or the path to your own file.""" whisper_batch_size: int = 4 """Clips decoded at once. ``1`` turns batching off; 4 suits large-v3 on an 8 GB card.""" @@ -62,6 +79,13 @@ class Config: paragraph_gap: float = 2.0 """Silence in seconds between segments that starts a new paragraph.""" + polish: bool = True + """Tidy the recognised text before writing it: fullwidth punctuation for Chinese + sentences, one copy of a looped phrase, and the vocabulary's rewrites.""" + + convert_to_simplified: bool = False + """Rewrite traditional characters as simplified ones. Needs 'uv sync --extra zh'.""" + # --- Google Drive (optional) ----------------------------------------- drive_upload: bool = False """Also copy every finished script into Google Drive. Local files are written either way.""" @@ -152,6 +176,11 @@ def validate(self) -> None: raise ConfigError("check_limit must be at least 1") if self.whisper_batch_size < 1: raise ConfigError("whisper_batch_size must be at least 1 (1 turns batching off)") + # Reading the file now means a typo fails the command, not the first video. + try: + load_vocabulary(self.vocabulary) + except VocabularyError as exc: + raise ConfigError(str(exc)) from exc if self.include_members_only and not (self.cookies_file or self.cookies_from_browser): raise ConfigError( "include_members_only needs a signed-in session: set cookies_file or " @@ -286,10 +315,26 @@ def load_config( whisper_compute_type = "float16" # Whisper transcribes Mandarin into traditional characters about as readily as -# simplified. A simplified-character seed sentence settles it. Change or clear -# this if you change `language`. +# simplified. A simplified-character seed sentence settles it. Leave it unset to +# get the stock sentence for `language`, or set it to "" for no seed at all. whisper_initial_prompt = "以下是普通话的句子。" +# The seed is followed by the video's own title and description, so the tickers +# and names that episode is about are words the model already expects. +prompt_from_metadata = true + +# Terms the channel says every episode, plus rewrites for the ones the model +# keeps getting wrong ("对中基金 => 对冲基金"). "zh-finance" ships with ytscript +# and suits a Mandarin US-market channel; point this at your own file to extend +# it — `python -c "import ytscript.vocabulary as v; print(v.DATA_DIR)"` finds the +# built-in to copy. Leave it out for no vocabulary. +vocabulary = "zh-finance" + +# Whisper's own default feeds each clip the text of the one before it, which is +# what makes it repeat a phrase for a minute when the audio goes quiet. Off also +# matches what batched decoding does, so both paths sound the same. +whisper_condition_on_previous_text = false + # Clips decoded at once — several times faster than one at a time, at the cost # of VRAM. 4 leaves headroom on an 8 GB card; a 12 GB or larger card can take # 8 or 16. Drop to 1 to turn batching off. @@ -299,6 +344,13 @@ def load_config( output_formats = ["txt"] timestamps = false +# Tidy the text before writing: fullwidth punctuation for Chinese sentences, one +# copy of a phrase the model looped on, and the vocabulary's rewrites. +polish = true + +# Rewrite traditional characters as simplified. Needs `uv sync --extra zh`. +convert_to_simplified = false + state_file = ".ytscript-state.json" keep_audio = false diff --git a/src/ytscript/data/zh-finance.txt b/src/ytscript/data/zh-finance.txt new file mode 100644 index 0000000..b1eb947 --- /dev/null +++ b/src/ytscript/data/zh-finance.txt @@ -0,0 +1,157 @@ +# Vocabulary for a Mandarin US-market commentary channel. +# +# Two kinds of line: +# term seeds Whisper's prompt so it expects the word +# wrong => right also rewrites the word when it comes out mangled +# +# Both sides of a "=>" line are seeded, so a correction doubles as a term. +# Rewrites are literal. An ASCII left-hand side only matches as a whole word, +# so "CDS" never fires inside "CDSX". Order matters: the prompt is capped at +# a couple of hundred characters, and the file is read top down. +# +# Copy this file, edit it, and point `vocabulary` at your copy — the errors a +# model makes follow the speaker, so this list is worth growing as you read. + +# --- terms this channel says every episode --------------------------------- +美股 +标普500 +纳斯达克 +费城半导体 +对中基金 => 对冲基金 +共同基金 +散户 +机构 + +# --- terms Whisper reliably gets wrong on this channel ---------------------- +# 费半 is the channel's shorthand for the Philadelphia semiconductor index. +飞班 => 费半 +肺瓣 => 费半 +非诚半导体 => 费城半导体 +自然负债表 => 资产负债表 +美股收益 => 每股收益 +非GNAP => 非GAAP +阿斯曼 => 阿斯麦 +极客讯霍尔 => 杰克逊霍尔 +杰克森霍尔 => 杰克逊霍尔 +自然力的推演 => 时间的推演 +S&DK => SNDK +CTS => CDS +ISIG => ISRG + +# --- tickers --------------------------------------------------------------- +NVDA +AVGO +TSLA +MSFT +GOOG +AMZN +META +AAPL +AMD +TSM +ASML +MU +SNDK +NFLX +ISRG +ORCL +PLTR +CRM +INTU +ADBE +APP +NOW +IBM +INTC +QCOM +MRVL +LRCX +AMAT +UNH +LLY +COST +WMT +HD +BA +AXP +BRK +CRWD +PANW +DDOG +NET +UBER +SHOP +SNOW +RKLB +CRWV +NBIS +TER +GLW +TXN + +# --- indices and funds ------------------------------------------------------ +SPY +QQQ +SOXX +SMH +IGV +XLK +XLB +XLY +VIX +ETF +CTA + +# --- company names as they are said in Mandarin ----------------------------- +英伟达 +博通 +特斯拉 +微软 +谷歌 +亚马逊 +苹果 +台积电 +阿斯麦 +美光 +甲骨文 +英特尔 +高通 +奈飞 +达芬奇手术系统 + +# --- the jargon of the show -------------------------------------------------- +每股收益 +资产负债表 +自由现金流 +资本支出 +前瞻估值 +毛利率 +净利润 +财报 +期权 +顶背离 +轧空 +杠杆 +去杠杆 +踩踏 +回撤 +均线 +支撑 +压力位 +仓位 +加仓 +减仓 +建仓 +美联储 +国债收益率 +财政部 +回购 +解禁 +信用违约互换 +CDS +RSI +EPS +GAAP +PMI +13F +杰克逊霍尔 diff --git a/src/ytscript/pipeline.py b/src/ytscript/pipeline.py index a6e1e66..9a121a3 100644 --- a/src/ytscript/pipeline.py +++ b/src/ytscript/pipeline.py @@ -12,8 +12,10 @@ from .drive import DriveError, DriveFile, DriveUploader from .formatting import write_outputs from .models import RunReport, Transcript, Video +from .polish import polish_segments from .state import State from .transcribers import Transcriber, TranscriptionError, build_transcriber +from .vocabulary import load_vocabulary, seed_prompt from .youtube import YouTubeClient, YouTubeError log = logging.getLogger("ytscript") @@ -57,6 +59,14 @@ def __init__( ) self._transcriber = transcriber self._uploader = uploader + self.vocabulary = load_vocabulary(config.vocabulary) + # An unset prompt means "the stock sentence for this language"; an empty + # one means the caller wants no seed at all. + self.seed = ( + seed_prompt(config.language) + if config.whisper_initial_prompt is None + else config.whisper_initial_prompt + ) @property def transcriber(self) -> Transcriber: @@ -72,6 +82,11 @@ def uploader(self) -> DriveUploader: self._uploader = DriveUploader.from_config(self.config) return self._uploader + def prompt_for(self, video: Video) -> str | None: + """The priming text this video is transcribed with.""" + subject = video if self.config.prompt_from_metadata else None + return self.vocabulary.prompt(subject, seed=self.seed) + def list_videos(self, limit: int | None = None) -> list[Video]: count = limit if limit is not None else self.config.check_limit return self.client.latest_videos(self.config.channel, count) @@ -140,9 +155,21 @@ def _process(self, video: Video, audio_dir: Path) -> tuple[list[Path], list[Driv config = self.config audio_path, video = self.client.download_audio(video, audio_dir) try: - segments, language = self.transcriber.transcribe(audio_path, config.language) + segments, language = self.transcriber.transcribe( + audio_path, config.language, prompt=self.prompt_for(video) + ) if not segments: raise TranscriptionError("no speech was recognised in the audio") + if config.polish or config.convert_to_simplified: + segments = polish_segments( + segments, + vocabulary=self.vocabulary if config.polish else None, + punctuation=config.polish, + simplified=config.convert_to_simplified, + loops=config.polish, + ) + if not segments: + raise TranscriptionError("every recognised segment was dropped as boilerplate") transcript = Transcript( video=video, segments=segments, diff --git a/src/ytscript/polish.py b/src/ytscript/polish.py new file mode 100644 index 0000000..98ed5d8 --- /dev/null +++ b/src/ytscript/polish.py @@ -0,0 +1,137 @@ +"""Cleaning up what the model returns, before it is written out as a script. + +Four things, all of them things Whisper does to Mandarin audio and none of them +things it can be talked out of with a prompt: it punctuates Chinese with ASCII +commas half the time, it loops a phrase when the audio underneath goes quiet, it +sometimes signs off with subtitle-credit boilerplate it learned from its training +data, and it spells domain words the way a general model would. +""" + +from __future__ import annotations + +import logging +import re +from typing import Any + +from .models import Segment +from .vocabulary import Vocabulary + +log = logging.getLogger("ytscript") + +# Han, kana and the fullwidth punctuation that goes with them. +_CJK = r" -〿぀-ヿ㐀-䶿一-鿿＀-・" + +_FULLWIDTH = {",": ",", "?": "?", "!": "!", ";": ";", ":": ":"} +_ASCII_PUNCTUATION = "".join(re.escape(mark) for mark in _FULLWIDTH) +# Only where a Chinese character is on one side or the other: "1,000" and the +# ":" in a URL keep the ASCII marks they need. +_AFTER_CJK = re.compile(rf"(?<=[{_CJK}])([{_ASCII_PUNCTUATION}])") +_BEFORE_CJK = re.compile(rf"([{_ASCII_PUNCTUATION}])(?=[{_CJK}])") + +# Credit lines Whisper has been known to invent over music or silence. None of +# these showed up in the scripts this list was written against — it is a guard, +# not a fix, and a segment only goes if that is the whole of it. +_BOILERPLATE = tuple( + re.compile(pattern) + for pattern in ( + r"^字幕(由|志願者|志愿者).{0,30}(提供|製作|制作)[。.]?$", + r"^.{0,20}(Amara\.org).{0,20}$", + r"^請不吝點贊.*$", + r"^请不吝点赞.*$", + r"^(明鏡與點點欄目|明镜与点点栏目)[。.]?$", + ) +) + +# How many identical segments in a row it takes to call it a loop rather than +# someone repeating themselves for effect. +LOOP_LENGTH = 3 + + +def normalize_punctuation(text: str) -> str: + """Give Chinese sentences the fullwidth marks they are written with.""" + text = _AFTER_CJK.sub(lambda m: _FULLWIDTH[m.group(1)], text) + return _BEFORE_CJK.sub(lambda m: _FULLWIDTH[m.group(1)], text) + + +def is_boilerplate(text: str) -> bool: + stripped = text.strip() + return any(pattern.match(stripped) for pattern in _BOILERPLATE) + + +def collapse_loops(segments: list[Segment], length: int = LOOP_LENGTH) -> list[Segment]: + """Keep one segment out of a run of ``length`` or more identical ones.""" + kept: list[Segment] = [] + run: list[Segment] = [] + + def flush() -> None: + if not run: + return + if len(run) >= length: + # The loop ran from the first repeat to the end of the last one. + kept.append(Segment(start=run[0].start, end=run[-1].end, text=run[0].text)) + log.info("dropped %d repeats of %r", len(run) - 1, run[0].text[:30]) + else: + kept.extend(run) + run.clear() + + for segment in segments: + if run and segment.text.strip() == run[-1].text.strip(): + run.append(segment) + continue + flush() + run.append(segment) + flush() + return kept + + +def _simplifier() -> Any | None: + try: + from opencc import OpenCC # noqa: PLC0415 - optional dependency + except ImportError: + log.warning( + "convert_to_simplified is on but opencc is not installed; leaving the " + "characters as the model wrote them. Install it with 'uv sync --extra zh'" + ) + return None + return OpenCC("t2s") + + +def polish_segments( + segments: list[Segment], + vocabulary: Vocabulary | None = None, + punctuation: bool = True, + simplified: bool = False, + loops: bool = True, +) -> list[Segment]: + """Run the whole clean-up over a transcript's segments.""" + convert = _simplifier() if simplified else None + polished: list[Segment] = [] + for segment in segments: + text = segment.text.strip() + if not text or is_boilerplate(text): + continue + if convert is not None: + text = convert.convert(text) + if vocabulary is not None: + text = vocabulary.correct(text) + if punctuation: + text = normalize_punctuation(text) + polished.append(Segment(start=segment.start, end=segment.end, text=text)) + return collapse_loops(polished) if loops else polished + + +def polish_text( + text: str, + vocabulary: Vocabulary | None = None, + punctuation: bool = True, + simplified: bool = False, +) -> str: + """The same clean-up over a script that has already been written to disk.""" + convert = _simplifier() if simplified else None + if convert is not None: + text = convert.convert(text) + if vocabulary is not None: + text = vocabulary.correct(text) + if punctuation: + text = normalize_punctuation(text) + return text diff --git a/src/ytscript/transcribers/__init__.py b/src/ytscript/transcribers/__init__.py index ab0bc95..e411b6c 100644 --- a/src/ytscript/transcribers/__init__.py +++ b/src/ytscript/transcribers/__init__.py @@ -25,10 +25,12 @@ def build_transcriber(config: Config) -> Transcriber: compute_type=config.whisper_compute_type, initial_prompt=config.whisper_initial_prompt, batch_size=config.whisper_batch_size, + condition_on_previous_text=config.whisper_condition_on_previous_text, ) if config.backend == "openai": return OpenAITranscriber( model=config.openai_model, api_key_env=config.openai_api_key_env, + initial_prompt=config.whisper_initial_prompt, ) raise TranscriptionError(f"unknown backend {config.backend!r}") diff --git a/src/ytscript/transcribers/base.py b/src/ytscript/transcribers/base.py index 0327c71..8aad2e6 100644 --- a/src/ytscript/transcribers/base.py +++ b/src/ytscript/transcribers/base.py @@ -19,7 +19,11 @@ class Transcriber(Protocol): name: str def transcribe( - self, audio_path: Path, language: str | None = None + self, audio_path: Path, language: str | None = None, prompt: str | None = None ) -> tuple[list[Segment], str | None]: - """Return the segments and the language that was used or detected.""" + """Return the segments and the language that was used or detected. + + ``prompt`` primes the model with words to expect — the video's title and + the channel's vocabulary — and overrides whatever the backend was built with. + """ ... diff --git a/src/ytscript/transcribers/faster_whisper.py b/src/ytscript/transcribers/faster_whisper.py index 1fd3762..0e04e4b 100644 --- a/src/ytscript/transcribers/faster_whisper.py +++ b/src/ytscript/transcribers/faster_whisper.py @@ -31,6 +31,7 @@ def __init__( beam_size: int = 5, initial_prompt: str | None = None, batch_size: int = 4, + condition_on_previous_text: bool = False, ) -> None: self.model_name = model self.device = device @@ -39,6 +40,7 @@ def __init__( self.beam_size = beam_size self.initial_prompt = initial_prompt self.batch_size = batch_size + self.condition_on_previous_text = condition_on_previous_text self._model: Any = None self._batched: Any = None @@ -79,14 +81,19 @@ def _wrap_batched(self, model: Any) -> Any: return BatchedInferencePipeline(model=model) def _run( - self, engine: Any, audio_path: Path, language: str | None, **extra: Any + self, + engine: Any, + audio_path: Path, + language: str | None, + prompt: str | None, + **extra: Any, ) -> tuple[list[Segment], Any]: raw_segments, info = engine.transcribe( str(audio_path), language=language, beam_size=self.beam_size, vad_filter=self.vad_filter, - initial_prompt=self.initial_prompt, + initial_prompt=prompt or self.initial_prompt, **extra, ) # The segment iterator is lazy — this is where the decoding actually runs, @@ -98,16 +105,19 @@ def _run( return segments, info def transcribe( - self, audio_path: Path, language: str | None = None + self, audio_path: Path, language: str | None = None, prompt: str | None = None ) -> tuple[list[Segment], str | None]: model = self._load_model() + # Only the sequential engine takes this: batched decoding has no previous + # clip to condition on, so it behaves as if the flag were off. + sequential = {"condition_on_previous_text": self.condition_on_previous_text} try: if self._batched is None: - segments, info = self._run(model, audio_path, language) + segments, info = self._run(model, audio_path, language, prompt, **sequential) else: try: segments, info = self._run( - self._batched, audio_path, language, batch_size=self.batch_size + self._batched, audio_path, language, prompt, batch_size=self.batch_size ) except Exception as exc: if not _is_out_of_memory(exc): @@ -120,7 +130,7 @@ def transcribe( audio_path.name, self.batch_size, ) - segments, info = self._run(model, audio_path, language) + segments, info = self._run(model, audio_path, language, prompt, **sequential) except Exception as exc: raise TranscriptionError(f"transcription of {audio_path.name} failed: {exc}") from exc detected = language or getattr(info, "language", None) diff --git a/src/ytscript/transcribers/openai_api.py b/src/ytscript/transcribers/openai_api.py index b354552..cc2d2f3 100644 --- a/src/ytscript/transcribers/openai_api.py +++ b/src/ytscript/transcribers/openai_api.py @@ -23,10 +23,12 @@ def __init__( model: str = "whisper-1", api_key_env: str = "OPENAI_API_KEY", api_key: str | None = None, + initial_prompt: str | None = None, ) -> None: self.model = model self.api_key_env = api_key_env self._api_key = api_key + self.initial_prompt = initial_prompt self._client: Any = None def _load_client(self) -> Any: @@ -46,7 +48,7 @@ def _load_client(self) -> Any: return self._client def transcribe( - self, audio_path: Path, language: str | None = None + self, audio_path: Path, language: str | None = None, prompt: str | None = None ) -> tuple[list[Segment], str | None]: client = self._load_client() size = audio_path.stat().st_size @@ -63,6 +65,9 @@ def transcribe( } if language: kwargs["language"] = language + # The endpoint takes the same kind of priming text as the local model. + if prompt or self.initial_prompt: + kwargs["prompt"] = prompt or self.initial_prompt try: with audio_path.open("rb") as handle: response = client.audio.transcriptions.create(file=handle, **kwargs) diff --git a/src/ytscript/vocabulary.py b/src/ytscript/vocabulary.py new file mode 100644 index 0000000..a282fca --- /dev/null +++ b/src/ytscript/vocabulary.py @@ -0,0 +1,177 @@ +"""Domain vocabulary: what the model should expect, and what it keeps getting wrong. + +Whisper decodes a word it has been primed for far more reliably than one it has +not, and it accepts a short prompt to be primed with. Two things go in there: the +video's own title and description, which name the day's subject, and a glossary of +terms the channel says every episode. Whatever still comes out wrong is rewritten +afterwards from the same file. +""" + +from __future__ import annotations + +import re +from dataclasses import dataclass +from pathlib import Path + +from .models import Video + +DATA_DIR = Path(__file__).parent / "data" + +# Whisper conditions on at most 224 tokens of prompt and drops the rest without +# saying so. A Chinese character is roughly a token, so this is the safe ceiling. +MAX_PROMPT_CHARS = 200 + +_SEEDS = { + # A simplified-character sentence settles which script the output is written in. + "zh": "以下是普通话的句子。", +} + +_ASCII = re.compile(r"^[\w&.\- ]+$", re.ASCII) +_TERM_SEPARATOR = "、" +_ARROW = re.compile(r"\s*(?:=>|->)\s*") + + +class VocabularyError(ValueError): + """Raised when a vocabulary file is missing or a line cannot be read.""" + + +def seed_prompt(language: str | None) -> str | None: + """The stock priming sentence for a language, if there is one.""" + if not language: + return None + return _SEEDS.get(language.split("-")[0].lower()) + + +@dataclass(frozen=True) +class Correction: + """One ``wrong => right`` rewrite.""" + + wrong: str + right: str + pattern: re.Pattern[str] + + @classmethod + def build(cls, wrong: str, right: str) -> Correction: + body = re.escape(wrong) + if _ASCII.match(wrong): + # "CDS" should not fire inside "CDSX"; Chinese has no such boundary, and + # \b is no use here because Python counts Han characters as word ones. + body = rf"(? str: + # A plain replacement: the right-hand side is a word, not a regex template. + return self.pattern.sub(lambda _: self.right, text) + + +@dataclass(frozen=True) +class Vocabulary: + """Terms to prime the model with, and rewrites for what it still gets wrong.""" + + terms: tuple[str, ...] = () + corrections: tuple[Correction, ...] = () + source: str = "" + + def __bool__(self) -> bool: + return bool(self.terms or self.corrections) + + def correct(self, text: str) -> str: + for correction in self.corrections: + text = correction.apply(text) + return text + + def prompt( + self, + video: Video | None = None, + seed: str | None = None, + max_chars: int = MAX_PROMPT_CHARS, + ) -> str | None: + """Build the priming text: seed sentence, this video's subject, then terms. + + Terms named in the title or description come first — they are the ones the + episode actually says — and the rest fill whatever budget is left. + """ + pieces: list[str] = [] + if seed: + pieces.append(seed.strip()) + subject = _subject(video) + if subject: + pieces.append(subject) + + used = sum(len(piece) for piece in pieces) + haystack = (subject or "").lower() + ranked = sorted(self.terms, key=lambda term: term.lower() not in haystack) + chosen: list[str] = [] + for term in ranked: + cost = len(term) + len(_TERM_SEPARATOR) + if used + cost > max_chars: + continue + chosen.append(term) + used += cost + if chosen: + pieces.append(_TERM_SEPARATOR.join(chosen) + "。") + + prompt = "".join(pieces).strip() + return prompt[:max_chars] or None + + +def _subject(video: Video | None) -> str: + """The video's own words about itself: title first, then a slice of the blurb.""" + if video is None: + return "" + parts = [video.title.strip()] if video.title else [] + description = (video.description or "").strip() + if description: + # Descriptions run to link dumps and boilerplate; the opening line is the + # part that says what the episode is about. + first = description.splitlines()[0].strip() + if first: + parts.append(first) + return " ".join(parts) + + +def parse_vocabulary(text: str, source: str = "") -> Vocabulary: + """Read the ``term`` / ``wrong => right`` lines of a vocabulary file.""" + terms: list[str] = [] + corrections: list[Correction] = [] + seen: set[str] = set() + + def remember(term: str) -> None: + if term and term not in seen: + seen.add(term) + terms.append(term) + + for number, raw in enumerate(text.splitlines(), start=1): + line = raw.split("#", 1)[0].strip() + if not line: + continue + if _ARROW.search(line): + wrong, right = (part.strip() for part in _ARROW.split(line, maxsplit=1)) + if not wrong or not right: + raise VocabularyError( + f"{source or 'vocabulary'} line {number}: expected 'wrong => right'" + ) + corrections.append(Correction.build(wrong, right)) + # What the correction produces is also what the model should expect. + remember(right) + continue + remember(line) + + return Vocabulary(terms=tuple(terms), corrections=tuple(corrections), source=source) + + +def builtin_names() -> list[str]: + return sorted(path.stem for path in DATA_DIR.glob("*.txt")) + + +def load_vocabulary(name: str | Path | None) -> Vocabulary: + """Load a built-in vocabulary by name, or a file by path. ``None`` loads nothing.""" + if name is None or name == "": + return Vocabulary() + path = DATA_DIR / f"{name}.txt" if isinstance(name, str) and "/" not in name else Path(name) + if not path.is_file(): + raise VocabularyError( + f"no vocabulary {str(name)!r}: expected a file path or one of " + f"{', '.join(builtin_names())}" + ) + return parse_vocabulary(path.read_text(encoding="utf-8"), source=str(path)) diff --git a/tests/fakes.py b/tests/fakes.py index 1ade5a7..d0cdb18 100644 --- a/tests/fakes.py +++ b/tests/fakes.py @@ -46,24 +46,30 @@ def download_audio(self, video: Video, dest_dir: Path) -> tuple[Path, Video]: class FakeTranscriber: name = "fake" - def __init__(self, language: str | None = "en", fail_on: set[str] | None = None) -> None: + def __init__( + self, + language: str | None = "en", + fail_on: set[str] | None = None, + segments: list[Segment] | None = None, + ) -> None: self.language = language self.fail_on = fail_on or set() + self.segments = segments self.calls: list[tuple[Path, str | None]] = [] + self.prompts: list[str | None] = [] - def transcribe(self, audio_path: Path, language: str | None = None): + def transcribe(self, audio_path: Path, language: str | None = None, prompt: str | None = None): self.calls.append((audio_path, language)) + self.prompts.append(prompt) if audio_path.stem in self.fail_on: from ytscript.transcribers import TranscriptionError raise TranscriptionError("backend exploded") - return ( - [ - Segment(0.0, 2.0, "Hello there."), - Segment(6.0, 8.0, "Second paragraph starts here."), - ], - language or self.language, - ) + segments = self.segments or [ + Segment(0.0, 2.0, "Hello there."), + Segment(6.0, 8.0, "Second paragraph starts here."), + ] + return list(segments), language or self.language class FakeDriveUploader: diff --git a/tests/test_cli.py b/tests/test_cli.py index 27f4c0c..7fc79b4 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -236,3 +236,34 @@ def test_drive_auth_reports_where_the_scripts_will_go( out = capsys.readouterr().out assert "token is cached in" in out assert "folder-id" in out + + +def test_polish_rewrites_scripts_that_are_already_on_disk( + project: Path, capsys: pytest.CaptureFixture +) -> None: + scripts = project / "scripts" + scripts.mkdir() + (scripts / "one.txt").write_text("飞班收跌,收盘", encoding="utf-8") + (scripts / "two.txt").write_text("没有问题。", encoding="utf-8") + (scripts / "notes.csv").write_text("飞班", encoding="utf-8") + + assert cli.main(["polish", "scripts", "--vocabulary", "zh-finance", "--dry-run"]) == 0 + out = capsys.readouterr().out + assert "would rewrite 1 of 2 file(s)" in out + assert (scripts / "one.txt").read_text(encoding="utf-8") == "飞班收跌,收盘" + + assert cli.main(["polish", "scripts", "--vocabulary", "zh-finance"]) == 0 + assert "rewrote 1 of 2 file(s)" in capsys.readouterr().out + assert (scripts / "one.txt").read_text(encoding="utf-8") == "费半收跌,收盘" + # A file that was already clean is not rewritten, and neither is the CSV. + assert (scripts / "notes.csv").read_text(encoding="utf-8") == "飞班" + + +def test_polish_says_when_a_path_is_missing(project: Path, capsys: pytest.CaptureFixture) -> None: + assert cli.main(["polish", "nowhere"]) == 1 + assert "no such file or directory" in capsys.readouterr().err + + +def test_an_unknown_vocabulary_stops_the_run(project: Path, capsys: pytest.CaptureFixture) -> None: + assert cli.main(["run", "--vocabulary", "nope"]) == 1 + assert "no vocabulary 'nope'" in capsys.readouterr().err diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index df322e1..7fadecb 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -6,6 +6,7 @@ from fakes import FakeTranscriber, FakeYouTubeClient, make_videos from ytscript.config import Config +from ytscript.models import Segment from ytscript.pipeline import Pipeline, select_videos from ytscript.state import State @@ -189,3 +190,44 @@ def test_members_only_videos_are_taken_when_signed_in(tmp_path: Path) -> None: assert report.members_only == [] assert sorted(client.downloaded) == ["vid000", "vid001", "vid002"] + + +def test_the_video_primes_the_model_with_its_own_title(tmp_path: Path) -> None: + videos = make_videos(1) + _, _, transcriber = build(tmp_path, videos) + pipeline, _, transcriber = build(tmp_path, videos, language="zh") + pipeline.run() + assert transcriber.prompts == ["以下是普通话的句子。Episode 0"] + + +def test_the_vocabulary_joins_the_prompt_and_fixes_the_text(tmp_path: Path) -> None: + glossary = tmp_path / "glossary.txt" + glossary.write_text("对中基金 => 对冲基金\n", encoding="utf-8") + client = FakeYouTubeClient(make_videos(1)) + transcriber = FakeTranscriber( + language="zh", segments=[Segment(0.0, 2.0, "对中基金的仓位,很重")] + ) + config = make_config(tmp_path, language="zh", vocabulary=str(glossary)) + pipeline = Pipeline(config, client=client, transcriber=transcriber) + pipeline.run() + + assert "对冲基金" in (transcriber.prompts[0] or "") + written = next((tmp_path / "scripts").glob("*.txt")).read_text(encoding="utf-8") + assert "对冲基金的仓位,很重" in written + + +def test_metadata_priming_can_be_turned_off(tmp_path: Path) -> None: + pipeline, _, transcriber = build( + tmp_path, make_videos(1), language="zh", prompt_from_metadata=False + ) + pipeline.run() + assert transcriber.prompts == ["以下是普通话的句子。"] + + +def test_polish_off_leaves_the_text_exactly_as_recognised(tmp_path: Path) -> None: + client = FakeYouTubeClient(make_videos(1)) + transcriber = FakeTranscriber(language="zh", segments=[Segment(0.0, 2.0, "重仓,很重")]) + config = make_config(tmp_path, language="zh", polish=False) + Pipeline(config, client=client, transcriber=transcriber).run() + written = next((tmp_path / "scripts").glob("*.txt")).read_text(encoding="utf-8") + assert "重仓,很重" in written diff --git a/tests/test_polish.py b/tests/test_polish.py new file mode 100644 index 0000000..5dc5062 --- /dev/null +++ b/tests/test_polish.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +from ytscript.models import Segment +from ytscript.polish import ( + collapse_loops, + is_boilerplate, + normalize_punctuation, + polish_segments, + polish_text, +) +from ytscript.vocabulary import parse_vocabulary + + +def test_chinese_sentences_get_fullwidth_punctuation() -> None: + assert normalize_punctuation("大家好,欢迎回来!今天呢?") == "大家好,欢迎回来!今天呢?" + + +def test_punctuation_that_is_not_chinese_is_left_alone() -> None: + # Thousands separators, clock times and URLs all use the ASCII marks. + assert normalize_punctuation("营收 1,250 亿") == "营收 1,250 亿" + assert normalize_punctuation("美东时间 7:35") == "美东时间 7:35" + assert normalize_punctuation("https://example.com/a,b") == "https://example.com/a,b" + + +def test_punctuation_converts_between_a_ticker_and_chinese() -> None: + assert normalize_punctuation("看 AVGO,今天放量") == "看 AVGO,今天放量" + + +def test_a_looped_phrase_is_kept_once() -> None: + segments = [Segment(float(i), float(i) + 1, "不要抢") for i in range(5)] + collapsed = collapse_loops([Segment(0.0, 1.0, "开始"), *segments]) + assert [s.text for s in collapsed] == ["开始", "不要抢"] + # The one that is kept spans the whole loop, so later timestamps still line up. + assert collapsed[1].start == 0.0 + assert collapsed[1].end == 5.0 + + +def test_saying_something_twice_is_not_a_loop() -> None: + segments = [Segment(0.0, 1.0, "不要抢"), Segment(1.0, 2.0, "不要抢")] + assert collapse_loops(segments) == segments + + +def test_subtitle_credits_are_boilerplate() -> None: + assert is_boilerplate("字幕由Amara.org社区提供") + assert is_boilerplate("请不吝点赞 订阅 转发") + assert not is_boilerplate("今天的字幕由我自己写") + + +def test_polish_segments_runs_the_lot() -> None: + vocabulary = parse_vocabulary("对中基金 => 对冲基金") + segments = [ + Segment(0.0, 1.0, " 对中基金的仓位,很重 "), + Segment(1.0, 2.0, "字幕由Amara.org社区提供"), + Segment(2.0, 3.0, ""), + ] + assert polish_segments(segments, vocabulary) == [Segment(0.0, 1.0, "对冲基金的仓位,很重")] + + +def test_polish_segments_can_be_narrowed_to_one_job() -> None: + segments = [Segment(0.0, 1.0, "对中基金,重仓")] + polished = polish_segments( + segments, parse_vocabulary("对中基金 => 对冲基金"), punctuation=False + ) + assert polished[0].text == "对冲基金,重仓" + + +def test_polish_text_rewrites_a_whole_script() -> None: + vocabulary = parse_vocabulary("飞班 => 费半") + assert polish_text("飞班收跌,收盘", vocabulary) == "费半收跌,收盘" + + +def test_simplified_conversion_is_skipped_when_opencc_is_missing(monkeypatch) -> None: + import builtins + + real_import = builtins.__import__ + + def no_opencc(name, *args, **kwargs): + if name == "opencc": + raise ImportError("no opencc here") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", no_opencc) + # The characters stay as they were rather than the run failing. + assert polish_text("這個", simplified=True) == "這個" diff --git a/tests/test_vocabulary.py b/tests/test_vocabulary.py new file mode 100644 index 0000000..faebe84 --- /dev/null +++ b/tests/test_vocabulary.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +from datetime import date +from pathlib import Path + +import pytest + +from ytscript.models import Video +from ytscript.vocabulary import ( + VocabularyError, + load_vocabulary, + parse_vocabulary, + seed_prompt, +) + + +def make_video(title: str = "美股 半导体 AVGO", description: str | None = None) -> Video: + return Video( + id="vid001", + title=title, + url="https://www.youtube.com/watch?v=vid001", + channel="视野环球财经", + upload_date=date(2026, 8, 22), + description=description, + ) + + +def test_plain_lines_are_terms_and_arrows_are_corrections() -> None: + vocabulary = parse_vocabulary("# a comment\n对中基金 => 对冲基金\nNVDA\n\n") + assert vocabulary.terms == ("对冲基金", "NVDA") + assert vocabulary.correct("今天对中基金的仓位") == "今天对冲基金的仓位" + + +def test_an_arrow_line_needs_both_sides() -> None: + with pytest.raises(VocabularyError): + parse_vocabulary("=> 对冲基金", source="glossary.txt") + + +def test_ascii_corrections_only_match_whole_words() -> None: + vocabulary = parse_vocabulary("CTS => CDS") + assert vocabulary.correct("它的CTS飙升") == "它的CDS飙升" + assert vocabulary.correct("CTSX不是缩写") == "CTSX不是缩写" + + +def test_chinese_corrections_match_inside_a_run_of_characters() -> None: + vocabulary = parse_vocabulary("飞班 => 费半") + assert vocabulary.correct("今天飞班收跌") == "今天费半收跌" + + +def test_prompt_leads_with_the_seed_and_the_video_title() -> None: + vocabulary = parse_vocabulary("NVDA\n对冲基金") + prompt = vocabulary.prompt(make_video(), seed="以下是普通话的句子。") + assert prompt is not None + assert prompt.startswith("以下是普通话的句子。美股 半导体 AVGO") + assert "对冲基金" in prompt + + +def test_prompt_takes_the_first_line_of_the_description() -> None: + video = make_video(description="今天聊 AVGO 的融资\n\n免责声明:...\nhttps://example.com") + prompt = parse_vocabulary("").prompt(video, seed=None) + assert prompt == "美股 半导体 AVGO 今天聊 AVGO 的融资" + + +def test_prompt_stays_inside_the_budget_and_prefers_terms_from_the_title() -> None: + vocabulary = parse_vocabulary("\n".join([f"填充词{i:02d}" for i in range(40)] + ["AVGO"])) + prompt = vocabulary.prompt(make_video(), seed="以下是普通话的句子。", max_chars=60) + assert prompt is not None + assert len(prompt) <= 60 + # AVGO is in the title, so it survives the budget the filler words do not. + assert "AVGO" in prompt.split("。")[-2] + + +def test_prompt_is_none_when_there_is_nothing_to_say() -> None: + assert parse_vocabulary("").prompt(None, seed=None) is None + + +def test_seed_prompt_is_language_specific() -> None: + assert seed_prompt("zh") == "以下是普通话的句子。" + assert seed_prompt("zh-CN") == "以下是普通话的句子。" + assert seed_prompt("en") is None + assert seed_prompt(None) is None + + +def test_load_vocabulary_finds_the_builtin_and_a_path(tmp_path: Path) -> None: + builtin = load_vocabulary("zh-finance") + assert "对冲基金" in builtin.terms + assert builtin.correct("自然负债表") == "资产负债表" + + path = tmp_path / "mine.txt" + path.write_text("蜂蜜 => 蜂蜜柠檬\n", encoding="utf-8") + assert load_vocabulary(path).correct("蜂蜜") == "蜂蜜柠檬" + + assert not load_vocabulary(None) + assert not load_vocabulary("") + + +def test_load_vocabulary_says_what_it_expected() -> None: + with pytest.raises(VocabularyError, match="zh-finance"): + load_vocabulary("zh-finanace") From 521e6e6d5bbf7a9282deca8ddd1873d7f3638667 Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 24 Aug 2026 14:37:35 +0000 Subject: [PATCH 2/3] Relock for the zh extra, and cover it in CI MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The zh extra went into pyproject.toml without a matching uv.lock, which `uv lock --check` catches in the lint job — local runs pass `--frozen` and skip it. Add the extra to the extras matrix too, so its lazily imported dependency is checked the way local, openai and drive are. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01GARmHk71niEh1TPunhH3Z8 --- .github/workflows/ci.yml | 3 +++ uv.lock | 15 ++++++++++++++- 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8fff56e..999800b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -91,6 +91,9 @@ jobs: - extra: drive # The connector imports the client lazily, so name it too. import: "ytscript.drive, googleapiclient.discovery, google_auth_oauthlib.flow, google_auth_httplib2, socks" + - extra: zh + # Imported lazily as well, and only when convert_to_simplified is on. + import: "ytscript.polish, opencc" steps: - uses: actions/checkout@v5 diff --git a/uv.lock b/uv.lock index b9297ee..1a8f406 100644 --- a/uv.lock +++ b/uv.lock @@ -1054,6 +1054,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/6a/db/2b7a1b3de659bb82aef979116c74e809982b13e42c057759767552b5155f/openai-3.3.1-py3-none-any.whl", hash = "sha256:9652df7fdf8ee6f5bd58e0a12f2b1d414a18e0f06bb7a9a57c8643a5f5469bd3", size = 1690337, upload-time = "2026-08-19T16:31:32.812Z" }, ] +[[package]] +name = "opencc-python-reimplemented" +version = "0.1.7" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/8d/6d/c6f37eed651dd6b752e50f80a93396cdaa42a6acc6ce05ad7452303ea511/opencc-python-reimplemented-0.1.7.tar.gz", hash = "sha256:4f777ea3461a25257a7b876112cfa90bb6acabc6dfb843bf4d11266e43579dee", size = 482566, upload-time = "2023-02-11T03:58:42.25Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/30/6b/055b7806f320cc8f2cdf23c5f70221c0dc1683fca9ffaf76dfc2ad4b91b6/opencc_python_reimplemented-0.1.7-py2.py3-none-any.whl", hash = "sha256:41b3b92943c7bed291f448e9c7fad4b577c8c2eae30fcfe5a74edf8818493aa6", size = 481813, upload-time = "2023-02-11T03:58:39.66Z" }, +] + [[package]] name = "packaging" version = "26.3" @@ -1585,6 +1594,9 @@ local = [ openai = [ { name = "openai" }, ] +zh = [ + { name = "opencc-python-reimplemented" }, +] [package.dev-dependencies] dev = [ @@ -1601,10 +1613,11 @@ requires-dist = [ { name = "google-auth-httplib2", marker = "extra == 'drive'", specifier = ">=0.2" }, { name = "google-auth-oauthlib", marker = "extra == 'drive'", specifier = ">=1.2" }, { name = "openai", marker = "extra == 'openai'", specifier = ">=1.30.0" }, + { name = "opencc-python-reimplemented", marker = "extra == 'zh'", specifier = ">=0.1.7" }, { name = "pysocks", marker = "extra == 'drive'", specifier = ">=1.7" }, { name = "yt-dlp", specifier = ">=2024.4.9" }, ] -provides-extras = ["local", "openai", "drive"] +provides-extras = ["local", "openai", "zh", "drive"] [package.metadata.requires-dev] dev = [ From 7b65bb5a5fc555a3fe530f6fbfd68d209a7850ac Mon Sep 17 00:00:00 2001 From: Claude Date: Mon, 24 Aug 2026 14:40:41 +0000 Subject: [PATCH 3/3] Recognise a vocabulary path on Windows MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit load_vocabulary told a built-in name from a file path by looking for a forward slash, which a Windows path does not contain: vocabulary = "C:\scripts\mine.txt" resolved to the built-in directory and failed with "no vocabulary". A built-in is now a bare stem — no directory separator of either kind, no suffix — and a file of that name is tried either way, so a built-in never shadows one. Both separators are checked on every platform, so the unit test covering it fails everywhere rather than only on the Windows runner. Co-Authored-By: Claude Opus 5 Claude-Session: https://claude.ai/code/session_01GARmHk71niEh1TPunhH3Z8 --- src/ytscript/vocabulary.py | 29 ++++++++++++++++++++++------- tests/test_vocabulary.py | 12 ++++++++++++ 2 files changed, 34 insertions(+), 7 deletions(-) diff --git a/src/ytscript/vocabulary.py b/src/ytscript/vocabulary.py index a282fca..86c0e69 100644 --- a/src/ytscript/vocabulary.py +++ b/src/ytscript/vocabulary.py @@ -164,14 +164,29 @@ def builtin_names() -> list[str]: return sorted(path.stem for path in DATA_DIR.glob("*.txt")) +def _is_builtin_name(name: str) -> bool: + """A built-in is a bare stem: no directory and no suffix, on any platform. + + Both separators are checked whatever the platform is running: a Windows path + has no forward slash in it, and treating ``C:\\scripts\\mine.txt`` as the name + of a built-in is how this went wrong once already. + """ + return not any(separator in name for separator in ("/", "\\")) and not Path(name).suffix + + def load_vocabulary(name: str | Path | None) -> Vocabulary: """Load a built-in vocabulary by name, or a file by path. ``None`` loads nothing.""" if name is None or name == "": return Vocabulary() - path = DATA_DIR / f"{name}.txt" if isinstance(name, str) and "/" not in name else Path(name) - if not path.is_file(): - raise VocabularyError( - f"no vocabulary {str(name)!r}: expected a file path or one of " - f"{', '.join(builtin_names())}" - ) - return parse_vocabulary(path.read_text(encoding="utf-8"), source=str(path)) + candidates: list[Path] = [] + if isinstance(name, str) and _is_builtin_name(name): + candidates.append(DATA_DIR / f"{name}.txt") + # A file of that name is still worth a look, so a built-in never shadows one. + candidates.append(Path(name)) + + for candidate in candidates: + if candidate.is_file(): + return parse_vocabulary(candidate.read_text(encoding="utf-8"), source=str(candidate)) + raise VocabularyError( + f"no vocabulary {str(name)!r}: expected a file path or one of {', '.join(builtin_names())}" + ) diff --git a/tests/test_vocabulary.py b/tests/test_vocabulary.py index faebe84..b8abd6c 100644 --- a/tests/test_vocabulary.py +++ b/tests/test_vocabulary.py @@ -8,6 +8,7 @@ from ytscript.models import Video from ytscript.vocabulary import ( VocabularyError, + _is_builtin_name, load_vocabulary, parse_vocabulary, seed_prompt, @@ -89,11 +90,22 @@ def test_load_vocabulary_finds_the_builtin_and_a_path(tmp_path: Path) -> None: path = tmp_path / "mine.txt" path.write_text("蜂蜜 => 蜂蜜柠檬\n", encoding="utf-8") assert load_vocabulary(path).correct("蜂蜜") == "蜂蜜柠檬" + # A path out of a config file arrives as a string, not a Path. + assert load_vocabulary(str(path)).correct("蜂蜜") == "蜂蜜柠檬" assert not load_vocabulary(None) assert not load_vocabulary("") +def test_only_a_bare_stem_names_a_builtin() -> None: + # Checked on every platform: a Windows path holds no forward slash, so + # looking for one alone read C:\scripts\mine.txt as the name of a built-in. + assert _is_builtin_name("zh-finance") + assert not _is_builtin_name(r"C:\scripts\mine.txt") + assert not _is_builtin_name("scripts/mine.txt") + assert not _is_builtin_name("mine.txt") + + def test_load_vocabulary_says_what_it_expected() -> None: with pytest.raises(VocabularyError, match="zh-finance"): load_vocabulary("zh-finanace")