Fix Whisper JSON chunk extraction
This commit is contained in:
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user