Fix Whisper JSON chunk extraction

This commit is contained in:
2026-07-29 15:34:01 +02:00
parent 1a6d731d21
commit 46565233d8
2 changed files with 375 additions and 0 deletions
@@ -0,0 +1,325 @@
#!/usr/bin/env python3
"""
Split a meeting transcript into reasonably sized chunks without cutting
through transcript blocks.
The script treats paragraphs separated by blank lines as atomic blocks.
It tries to cut close to --target-chars and will not exceed --max-chars
unless a single block is already larger than that limit.
Example:
python chunk_transcript.py meeting.txt
python chunk_transcript.py meeting.txt --target-chars 9000 --max-chars 11000
python chunk_transcript.py meeting.txt --overlap-blocks 1
"""
from __future__ import annotations
import argparse
import json
import re
import sys
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
TIMESTAMP_RE = re.compile(
r"(?m)^\s*(?:"
r"\[(?:\d{1,2}:)?\d{1,2}:\d{2}\]"
r"|(?:\d{1,2}:)?\d{1,2}:\d{2}\s*(?:-->|-)"
r")"
)
@dataclass(frozen=True)
class ChunkInfo:
number: int
filename: str
chars: int
blocks: int
first_timestamp: str | None
last_timestamp: str | None
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Split a transcript into block-aligned text chunks."
)
parser.add_argument(
"input_file",
type=Path,
help="Transcript text file or Whisper JSON file",
)
parser.add_argument(
"-o",
"--output-dir",
type=Path,
help="Output directory; default: <input-stem>_chunks",
)
parser.add_argument(
"--target-chars",
type=int,
default=9000,
help="Preferred chunk size in characters (default: 9000)",
)
parser.add_argument(
"--max-chars",
type=int,
default=11000,
help="Soft maximum chunk size in characters (default: 11000)",
)
parser.add_argument(
"--min-chars",
type=int,
default=5000,
help="Preferred minimum before a chunk may be closed (default: 5000)",
)
parser.add_argument(
"--overlap-blocks",
type=int,
default=0,
help="Repeat this many trailing blocks in the next chunk (default: 0)",
)
return parser.parse_args()
def validate_args(args: argparse.Namespace) -> None:
if args.target_chars <= 0 or args.max_chars <= 0 or args.min_chars < 0:
raise ValueError("Character limits must be positive.")
if args.min_chars > args.target_chars:
raise ValueError("--min-chars must not exceed --target-chars.")
if args.target_chars > args.max_chars:
raise ValueError("--target-chars must not exceed --max-chars.")
if args.overlap_blocks < 0:
raise ValueError("--overlap-blocks must not be negative.")
def normalize_text(text: str) -> str:
text = text.replace("\r\n", "\n").replace("\r", "\n")
return text.strip()
def blocks_from_whisper_json(data: Any) -> list[str]:
if not isinstance(data, dict):
raise ValueError("Whisper JSON input must contain a JSON object.")
segments = data.get("segments")
if not isinstance(segments, list):
text = data.get("text")
if isinstance(text, str) and text.strip():
return split_into_blocks(normalize_text(text))
raise ValueError("Whisper JSON input must contain 'segments' or 'text'.")
blocks: list[str] = []
for segment in segments:
if not isinstance(segment, dict):
continue
text = str(segment.get("text", "")).strip()
if text:
blocks.append(text)
if not blocks:
raise ValueError("Whisper JSON input contains no transcript text.")
return blocks
def read_transcript_blocks(input_file: Path) -> list[str]:
source = normalize_text(input_file.read_text(encoding="utf-8-sig"))
if not source:
raise ValueError("The transcript is empty.")
if input_file.suffix.lower() == ".json":
return blocks_from_whisper_json(json.loads(source))
return split_into_blocks(source)
def split_into_blocks(text: str) -> list[str]:
"""
Prefer blank-line-delimited transcript blocks.
If the file has no blank lines but contains line-start timestamps,
split before each timestamp. Otherwise, use non-empty lines as blocks.
"""
paragraphs = [part.strip() for part in re.split(r"\n\s*\n+", text) if part.strip()]
if len(paragraphs) > 1:
return paragraphs
timestamp_starts = list(TIMESTAMP_RE.finditer(text))
if len(timestamp_starts) > 1:
blocks: list[str] = []
for index, match in enumerate(timestamp_starts):
start = match.start()
end = (
timestamp_starts[index + 1].start()
if index + 1 < len(timestamp_starts)
else len(text)
)
block = text[start:end].strip()
if block:
blocks.append(block)
prefix = text[: timestamp_starts[0].start()].strip()
if prefix:
blocks.insert(0, prefix)
return blocks
return [line.strip() for line in text.splitlines() if line.strip()]
def rendered_length(blocks: list[str]) -> int:
if not blocks:
return 0
return sum(len(block) for block in blocks) + 2 * (len(blocks) - 1)
def build_chunks(
blocks: list[str],
target_chars: int,
max_chars: int,
min_chars: int,
overlap_blocks: int,
) -> list[list[str]]:
chunks: list[list[str]] = []
current: list[str] = []
for block in blocks:
prospective = current + [block]
prospective_len = rendered_length(prospective)
current_len = rendered_length(current)
should_close = (
bool(current)
and current_len >= min_chars
and (
prospective_len > max_chars
or (current_len >= target_chars and prospective_len > target_chars)
)
)
if should_close:
chunks.append(current)
overlap = current[-overlap_blocks:] if overlap_blocks else []
current = overlap + [block]
else:
current.append(block)
if current:
chunks.append(current)
return rebalance_last_chunk(chunks, min_chars, max_chars)
def rebalance_last_chunk(
chunks: list[list[str]], min_chars: int, max_chars: int
) -> list[list[str]]:
"""
Avoid a tiny final chunk by moving complete blocks from the previous chunk.
"""
if len(chunks) < 2 or rendered_length(chunks[-1]) >= min_chars:
return chunks
previous = chunks[-2]
final = chunks[-1]
while len(previous) > 1 and rendered_length(final) < min_chars:
candidate = previous[-1]
new_final = [candidate] + final
if rendered_length(new_final) > max_chars:
break
final.insert(0, previous.pop())
return chunks
def extract_timestamps(text: str) -> list[str]:
return [match.group(0).strip() for match in TIMESTAMP_RE.finditer(text)]
def write_chunks(
chunks: list[list[str]], output_dir: Path, input_name: str
) -> list[ChunkInfo]:
output_dir.mkdir(parents=True, exist_ok=True)
infos: list[ChunkInfo] = []
width = max(2, len(str(len(chunks))))
for number, blocks in enumerate(chunks, start=1):
content = "\n\n".join(blocks).strip() + "\n"
filename = f"chunk_{number:0{width}d}.txt"
path = output_dir / filename
path.write_text(content, encoding="utf-8")
timestamps = extract_timestamps(content)
infos.append(
ChunkInfo(
number=number,
filename=filename,
chars=len(content),
blocks=len(blocks),
first_timestamp=timestamps[0] if timestamps else None,
last_timestamp=timestamps[-1] if timestamps else None,
)
)
manifest = {
"source_file": input_name,
"chunk_count": len(infos),
"chunks": [asdict(info) for info in infos],
}
(output_dir / "manifest.json").write_text(
json.dumps(manifest, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
return infos
def main() -> int:
args = parse_args()
try:
validate_args(args)
if not args.input_file.is_file():
raise FileNotFoundError(f"Input file not found: {args.input_file}")
blocks = read_transcript_blocks(args.input_file)
chunks = build_chunks(
blocks=blocks,
target_chars=args.target_chars,
max_chars=args.max_chars,
min_chars=args.min_chars,
overlap_blocks=args.overlap_blocks,
)
output_dir = args.output_dir or args.input_file.with_name(
f"{args.input_file.stem}_chunks"
)
infos = write_chunks(chunks, output_dir, args.input_file.name)
except (OSError, UnicodeError, ValueError) as exc:
print(f"Error: {exc}", file=sys.stderr)
return 1
print(f"Source: {args.input_file}")
print(f"Blocks: {len(blocks)}")
print(f"Chunks: {len(infos)}")
print(f"Output: {output_dir}")
print()
for info in infos:
time_range = ""
if info.first_timestamp or info.last_timestamp:
time_range = (
f" | {info.first_timestamp or '?'} to {info.last_timestamp or '?'}"
)
print(
f"{info.filename}: {info.chars:>6} chars, "
f"{info.blocks:>3} blocks{time_range}"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+50
View File
@@ -0,0 +1,50 @@
import json
import tempfile
import unittest
from pathlib import Path
from src.meeting_lab.chunking.chunk_transcript import (
build_chunks,
read_transcript_blocks,
rendered_length,
)
class ChunkTranscriptTests(unittest.TestCase):
def test_whisper_json_uses_segments_not_aggregate_text(self) -> None:
data = {
"text": "alpha beta gamma",
"segments": [
{"text": "alpha"},
{"text": "beta"},
{"text": "gamma"},
],
}
with tempfile.TemporaryDirectory() as directory:
input_file = Path(directory) / "meeting.json"
input_file.write_text(json.dumps(data), encoding="utf-8")
blocks = read_transcript_blocks(input_file)
self.assertEqual(blocks, ["alpha", "beta", "gamma"])
self.assertNotIn("alpha beta gamma", blocks)
def test_chunks_do_not_share_later_blocks_without_overlap(self) -> None:
blocks = [f"block-{number:02d}-" + ("x" * 20) for number in range(10)]
chunks = build_chunks(
blocks=blocks,
target_chars=60,
max_chars=80,
min_chars=30,
overlap_blocks=0,
)
flattened = [block for chunk in chunks for block in chunk]
self.assertEqual(flattened, blocks)
self.assertEqual(len(flattened), len(set(flattened)))
self.assertTrue(all(rendered_length(chunk) <= 80 for chunk in chunks))
if __name__ == "__main__":
unittest.main()