Add initial topic segmentation prototype

This commit is contained in:
2026-07-21 16:48:19 +02:00
parent b2c195c365
commit f234efca48
2 changed files with 545 additions and 0 deletions
@@ -0,0 +1,545 @@
#!/usr/bin/env python3
"""
Segment a normalized meeting transcript into discussion topics using Ollama.
The script numbers the transcript blocks and asks the language model only to
identify the thematic structure of the discussion.
It does not extract facts, decisions, todos, questions or summaries.
Examples:
python src/meeting_lab/segmentation/segment_topics.py \
samples/chunks/chunk_01_normalized.txt
python src/meeting_lab/segmentation/segment_topics.py \
samples/chunks/chunk_01_normalized.txt \
--model qwen3:8b \
--output experiments/topic-segmentation/chunk_01_topics.json
"""
from __future__ import annotations
import argparse
import json
import re
import sys
import time
from pathlib import Path
from typing import Any
import requests
DEFAULT_MODEL = "qwen3:8b"
DEFAULT_ENDPOINT = "http://localhost:11434/api/generate"
SYSTEM_INSTRUCTION = """
Du analysierst die thematische Struktur eines Meeting-Transkripts.
Deine einzige Aufgabe ist die Themensegmentierung.
Du sollst erkennen:
1. wann ein Thema beginnt,
2. wann ein Thema endet,
3. wann zu einem anderen Thema gewechselt wird,
4. wann ein früheres Thema erneut aufgenommen wird.
Regeln:
1. Extrahiere keine Fakten.
2. Extrahiere keine Aufgaben.
3. Extrahiere keine Entscheidungen.
4. Extrahiere keine offenen Fragen.
5. Fasse die Diskussion nicht zusammen.
6. Bewerte die Aussagen nicht.
7. Erfinde keine Themen, die im Text nicht behandelt werden.
8. Verwende möglichst kurze und sachliche Themenbezeichnungen.
9. Jeder Block muss genau einem Thema zugeordnet werden.
10. Mehrere direkt aufeinanderfolgende Blöcke desselben Themas bilden ein Segment.
11. Wird ein früheres Thema später wieder aufgenommen, verwende dieselbe topic_id.
12. Kleine Rückfragen oder kurze Ergänzungen bleiben beim aktuellen Thema,
sofern sie kein eigenständiges Thema eröffnen.
13. Organisatorische Übergänge dürfen als eigenes Thema erfasst werden, wenn
sie einen erkennbaren Gesprächsabschnitt bilden.
14. Gib ausschließlich gültiges JSON aus.
15. Gib kein Markdown und keine Erläuterungen außerhalb des JSON aus.
""".strip()
OUTPUT_SCHEMA = {
"topics": [
{
"topic_id": "topic_001",
"title": "Kurze sachliche Themenbezeichnung",
"segments": [
{
"start_block": 1,
"end_block": 5,
}
],
}
]
}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Segment a meeting transcript into discussion topics."
)
parser.add_argument(
"input_file",
type=Path,
help="Normalized transcript or transcript chunk",
)
parser.add_argument(
"-o",
"--output",
type=Path,
help="Output JSON file; default: <input>_topics.json",
)
parser.add_argument(
"--model",
default=DEFAULT_MODEL,
help=f"Ollama model name (default: {DEFAULT_MODEL})",
)
parser.add_argument(
"--endpoint",
default=DEFAULT_ENDPOINT,
help=f"Ollama generate endpoint (default: {DEFAULT_ENDPOINT})",
)
parser.add_argument(
"--timeout",
type=int,
default=1800,
help="HTTP timeout in seconds (default: 1800)",
)
parser.add_argument(
"--temperature",
type=float,
default=0.0,
help="Sampling temperature (default: 0.0)",
)
return parser.parse_args()
def split_into_blocks(text: str) -> list[str]:
"""
Split a transcript into stable, numbered discussion blocks.
Blank-line-separated paragraphs are preferred. If the transcript contains
no blank lines, individual non-empty lines are used as blocks.
"""
normalized = text.replace("\r\n", "\n").replace("\r", "\n").strip()
paragraphs = [
part.strip()
for part in re.split(r"\n\s*\n+", normalized)
if part.strip()
]
if len(paragraphs) > 1:
return paragraphs
return [line.strip() for line in normalized.splitlines() if line.strip()]
def render_numbered_blocks(blocks: list[str]) -> str:
return "\n\n".join(
f"[BLOCK {number}]\n{block}"
for number, block in enumerate(blocks, start=1)
)
def build_prompt(source_name: str, blocks: list[str]) -> str:
schema_text = json.dumps(OUTPUT_SCHEMA, ensure_ascii=False, indent=2)
numbered_transcript = render_numbered_blocks(blocks)
return f"""{SYSTEM_INSTRUCTION}
Quelldatei:
{source_name}
Anzahl der Blöcke:
{len(blocks)}
Erwartete JSON-Struktur:
{schema_text}
Zusätzliche Anforderungen:
- Alle topic_id-Werte müssen dem Muster topic_001, topic_002 usw. folgen.
- Die Nummerierung beginnt mit topic_001.
- Die topic_id-Werte müssen innerhalb der Ausgabe eindeutig sein.
- start_block und end_block beziehen sich auf die angegebenen Blocknummern.
- start_block darf nicht größer als end_block sein.
- Kein Block darf fehlen.
- Kein Block darf mehreren Themen gleichzeitig zugeordnet werden.
- Direkt aufeinanderfolgende Bereiche desselben Themas sollen als ein Segment
ausgegeben werden.
- Nicht zusammenhängende Bereiche desselben Themas sollen als mehrere Segmente
innerhalb desselben Themas ausgegeben werden.
- Themenbezeichnungen sollen knapp sein und keine Zusammenfassung enthalten.
- Der Schlüssel "topics" muss immer vorhanden sein.
TRANSKRIPT:
--- BEGINN TRANSKRIPT ---
{numbered_transcript}
--- ENDE TRANSKRIPT ---
"""
def call_ollama(
endpoint: str,
model: str,
prompt: str,
timeout: int,
temperature: float,
) -> tuple[str, dict[str, Any]]:
payload = {
"model": model,
"prompt": prompt,
"stream": False,
"format": "json",
"options": {
"temperature": temperature,
},
}
started = time.perf_counter()
response = requests.post(
endpoint,
json=payload,
timeout=timeout,
)
elapsed = time.perf_counter() - started
response.raise_for_status()
data = response.json()
response_text = data.get("response")
if not isinstance(response_text, str) or not response_text.strip():
raise ValueError("Ollama returned no usable response text.")
metadata = {
"model": data.get("model", model),
"elapsed_seconds": round(elapsed, 3),
"total_duration_ns": data.get("total_duration"),
"load_duration_ns": data.get("load_duration"),
"prompt_eval_count": data.get("prompt_eval_count"),
"prompt_eval_duration_ns": data.get("prompt_eval_duration"),
"eval_count": data.get("eval_count"),
"eval_duration_ns": data.get("eval_duration"),
}
return response_text.strip(), metadata
def parse_json_response(text: str) -> dict[str, Any]:
try:
parsed = json.loads(text)
except json.JSONDecodeError:
# Fallback for models that wrap JSON in Markdown or add text around it.
match = re.search(r"\{.*\}", text, flags=re.DOTALL)
if not match:
raise
parsed = json.loads(match.group(0))
if not isinstance(parsed, dict):
raise ValueError(
"The model response is valid JSON but not a JSON object."
)
return parsed
def validate_topic_id(topic_id: Any) -> str:
if not isinstance(topic_id, str):
raise ValueError("Every topic must contain a string topic_id.")
if not re.fullmatch(r"topic_\d{3}", topic_id):
raise ValueError(
f"Invalid topic_id {topic_id!r}; expected format topic_001."
)
return topic_id
def validate_segment(
segment: Any,
topic_id: str,
block_count: int,
) -> dict[str, int]:
if not isinstance(segment, dict):
raise ValueError(
f"Topic {topic_id} contains a segment that is not an object."
)
start_block = segment.get("start_block")
end_block = segment.get("end_block")
if not isinstance(start_block, int) or isinstance(start_block, bool):
raise ValueError(
f"Topic {topic_id} contains an invalid start_block."
)
if not isinstance(end_block, int) or isinstance(end_block, bool):
raise ValueError(
f"Topic {topic_id} contains an invalid end_block."
)
if start_block < 1 or end_block > block_count:
raise ValueError(
f"Topic {topic_id} references blocks outside 1..{block_count}."
)
if start_block > end_block:
raise ValueError(
f"Topic {topic_id} contains start_block > end_block."
)
return {
"start_block": start_block,
"end_block": end_block,
}
def normalize_topics(
result: dict[str, Any],
block_count: int,
) -> list[dict[str, Any]]:
raw_topics = result.get("topics")
if not isinstance(raw_topics, list):
raise ValueError("Model output does not contain a topics list.")
topics: list[dict[str, Any]] = []
seen_topic_ids: set[str] = set()
for raw_topic in raw_topics:
if not isinstance(raw_topic, dict):
raise ValueError("Every topic must be a JSON object.")
topic_id = validate_topic_id(raw_topic.get("topic_id"))
if topic_id in seen_topic_ids:
raise ValueError(f"Duplicate topic_id: {topic_id}")
seen_topic_ids.add(topic_id)
title = raw_topic.get("title")
if not isinstance(title, str) or not title.strip():
raise ValueError(f"Topic {topic_id} has no usable title.")
raw_segments = raw_topic.get("segments")
if not isinstance(raw_segments, list) or not raw_segments:
raise ValueError(f"Topic {topic_id} has no segments.")
segments = [
validate_segment(segment, topic_id, block_count)
for segment in raw_segments
]
segments.sort(key=lambda item: item["start_block"])
topics.append(
{
"topic_id": topic_id,
"title": title.strip(),
"segments": segments,
}
)
if not topics:
raise ValueError("The model returned no topics.")
topics.sort(
key=lambda topic: min(
segment["start_block"] for segment in topic["segments"]
)
)
return topics
def validate_block_coverage(
topics: list[dict[str, Any]],
block_count: int,
) -> None:
assignments: dict[int, str] = {}
for topic in topics:
topic_id = topic["topic_id"]
for segment in topic["segments"]:
for block_number in range(
segment["start_block"],
segment["end_block"] + 1,
):
previous_topic = assignments.get(block_number)
if previous_topic is not None:
raise ValueError(
f"Block {block_number} is assigned to both "
f"{previous_topic} and {topic_id}."
)
assignments[block_number] = topic_id
missing_blocks = [
block_number
for block_number in range(1, block_count + 1)
if block_number not in assignments
]
if missing_blocks:
formatted = ", ".join(str(number) for number in missing_blocks)
raise ValueError(
f"The following blocks were not assigned to a topic: {formatted}"
)
def build_output(
source_name: str,
blocks: list[str],
topics: list[dict[str, Any]],
metadata: dict[str, Any],
) -> dict[str, Any]:
return {
"source": {
"file": source_name,
"block_count": len(blocks),
},
"topics": topics,
"_run": metadata,
}
def write_raw_response(
output_path: Path,
raw_text: str,
) -> Path:
raw_path = output_path.with_suffix(".raw.txt")
raw_path.write_text(raw_text + "\n", encoding="utf-8")
return raw_path
def main() -> int:
args = parse_args()
try:
if not args.input_file.is_file():
raise FileNotFoundError(
f"Input file not found: {args.input_file}"
)
transcript = args.input_file.read_text(
encoding="utf-8-sig"
).strip()
if not transcript:
raise ValueError("The input file is empty.")
blocks = split_into_blocks(transcript)
if not blocks:
raise ValueError("No discussion blocks could be created.")
output_path = args.output or args.input_file.with_name(
f"{args.input_file.stem}_topics.json"
)
output_path.parent.mkdir(parents=True, exist_ok=True)
prompt = build_prompt(
source_name=args.input_file.name,
blocks=blocks,
)
print(f"Model: {args.model}")
print(f"Input: {args.input_file}")
print(f"Blocks: {len(blocks)}")
print(f"Output: {output_path}")
print("Segmenting discussion topics ...")
raw_text, metadata = call_ollama(
endpoint=args.endpoint,
model=args.model,
prompt=prompt,
timeout=args.timeout,
temperature=args.temperature,
)
try:
parsed = parse_json_response(raw_text)
topics = normalize_topics(
result=parsed,
block_count=len(blocks),
)
validate_block_coverage(
topics=topics,
block_count=len(blocks),
)
except (json.JSONDecodeError, ValueError) as exc:
raw_path = write_raw_response(output_path, raw_text)
raise ValueError(
f"Model output could not be validated. "
f"Raw output saved to: {raw_path}. "
f"Reason: {exc}"
) from exc
result = build_output(
source_name=args.input_file.name,
blocks=blocks,
topics=topics,
metadata=metadata,
)
output_path.write_text(
json.dumps(
result,
ensure_ascii=False,
indent=2,
)
+ "\n",
encoding="utf-8",
)
segment_count = sum(
len(topic["segments"]) for topic in topics
)
resumed_topic_count = sum(
1 for topic in topics if len(topic["segments"]) > 1
)
print(f"Done in {metadata['elapsed_seconds']:.1f} seconds.")
print(f"Topics: {len(topics)}")
print(f"Segments: {segment_count}")
print(f"Resumed topics: {resumed_topic_count}")
return 0
except requests.ConnectionError:
print(
"Error: Ollama is not reachable. Is `ollama serve` running?",
file=sys.stderr,
)
return 1
except requests.Timeout:
print(
"Error: The Ollama request timed out.",
file=sys.stderr,
)
return 1
except requests.HTTPError as exc:
print(
f"Error: Ollama returned an HTTP error: {exc}",
file=sys.stderr,
)
return 1
except (OSError, UnicodeError, ValueError) as exc:
print(f"Error: {exc}", file=sys.stderr)
return 1
if __name__ == "__main__":
raise SystemExit(main())