File size: 7,013 Bytes
f1ef7e2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
"""
Batch sentence segmentation for banking, health, and telecom calls.

Reads every word-level transcription JSON from data/na_testset/results/
for the target domains and produces sentence_segments.json per call under:

  data/sentence_segments/
    banking/
      en_CA_Banking_1586889.json
      ...
    health/
      en_CA_Health_1587315.json
      ...
    telecom/
      en_CA_Telecom_1590675.json
      ...

Each output file has the same schema as frontend/src/sentence_segments.json
so the frontend and sentiment pipeline can consume any call directly.

Usage:  python batch_sentence_segments.py [--domains banking health telecom]
"""
import os
import json
import re
import argparse
from pathlib import Path

import paths

# ── config ────────────────────────────────────────────────────────────────────
RESULTS_ROOT = paths.NA_TESTSET / "results"
OUTPUT_ROOT  = paths.SENTENCE_SEG_ROOT

DEFAULT_DOMAINS = {"banking", "health", "telecom"}

MIN_DURATION_S = 1.5   # merge non-terminated fragments shorter than this
GAP_SPLIT_S    = 2.0   # silence gap that always starts a new sentence
MAX_MERGE_GAP  = 1.0   # only merge a short fragment if gap from prev ≀ this

ABBREVIATIONS = {
    "mr", "mrs", "ms", "dr", "prof", "sr", "jr", "vs", "st",
    "ave", "blvd", "dept", "est", "approx", "etc", "e.g", "i.e",
    "jan", "feb", "mar", "apr", "jun", "jul", "aug", "sep", "oct", "nov", "dec",
}
SENTENCE_END = re.compile(r"[.?!]$")


def is_sentence_end(word: str) -> bool:
    if not SENTENCE_END.search(word):
        return False
    base = re.sub(r"[.?!,;:]+$", "", word).lower()
    if base in ABBREVIATIONS:
        return False
    if re.match(r"^[a-z]$", base):
        return False
    return True


def words_to_sentences(words: list, speaker: str) -> list:
    if not words:
        return []

    raw = []
    bucket = [words[0]]
    for prev_w, curr_w in zip(words, words[1:]):
        gap = curr_w["start"] - prev_w["end"]
        if is_sentence_end(prev_w["word"]) or gap >= GAP_SPLIT_S:
            raw.append(bucket)
            bucket = [curr_w]
        else:
            bucket.append(curr_w)
    if bucket:
        raw.append(bucket)

    merged = []
    for sent_words in raw:
        start = sent_words[0]["start"]
        end   = sent_words[-1]["end"]
        dur   = end - start
        text  = " ".join(w["word"] for w in sent_words)
        last_word_terminal = SENTENCE_END.search(sent_words[-1]["word"])

        if merged and dur < MIN_DURATION_S and not last_word_terminal:
            gap_from_prev = start - merged[-1]["end"]
            if gap_from_prev <= MAX_MERGE_GAP:
                prev = merged[-1]
                prev["text"]     = prev["text"] + " " + text
                prev["end"]      = end
                prev["duration"] = round(prev["end"] - prev["start"], 3)
                continue

        merged.append({
            "speaker":  speaker,
            "start":    round(start, 3),
            "end":      round(end, 3),
            "duration": round(dur, 3),
            "text":     text,
        })

    for idx, s in enumerate(merged):
        s["id"] = f"{speaker}_{idx + 1:03d}"

    return merged


def interleave(agent_sents: list, customer_sents: list) -> list:
    combined = sorted(agent_sents + customer_sents, key=lambda s: s["start"])
    for i, s in enumerate(combined):
        s["seq_id"] = i + 1
    return combined


def segment_transcript(data: dict) -> dict:
    """Sentence-segment a loaded per-channel transcript dict (pure, no I/O).

    Shared with the single-call orchestrator so batch and live paths produce
    byte-identical segmentation from the same word-level transcript.
    """
    agent_sents    = words_to_sentences(data.get("agent", []),    "AGENT")
    customer_sents = words_to_sentences(data.get("customer", []), "CUSTOMER")
    all_sents      = interleave(agent_sents, customer_sents)

    return {
        "call_id":            data["call_id"],
        "accent":             data.get("accent"),
        "domain":             data.get("domain"),
        "model":              data.get("model"),
        "total":              len(all_sents),
        "agent_sentences":    len(agent_sents),
        "customer_sentences": len(customer_sents),
        "sentences":          all_sents,
    }


def segment_call(src_path: Path) -> dict:
    with open(src_path, encoding="utf-8") as f:
        data = json.load(f)
    return segment_transcript(data)


def collect_sources(domains: set) -> list[tuple[str, Path]]:
    """Return (domain_label, src_path) for every matching call."""
    sources = []
    for accent_dir in sorted(RESULTS_ROOT.iterdir()):
        if not accent_dir.is_dir():
            continue
        for json_file in sorted(accent_dir.glob("*.json")):
            stem = json_file.stem  # e.g. en_CA_Banking_1586889
            # domain is the part between accent and numeric id
            parts = stem.split("_")
            # find the domain token β€” the first purely alpha segment after the accent
            # accent = en_CA or en_US_General (variable length)
            # try matching known domains
            stem_lower = stem.lower()
            matched = None
            for d in domains:
                if f"_{d}_" in stem_lower or stem_lower.endswith(f"_{d}"):
                    matched = d
                    break
            if matched:
                sources.append((matched, json_file))
    return sources


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--domains", nargs="+",
                    default=sorted(DEFAULT_DOMAINS),
                    help="domains to segment (any manifest domain; "
                         "default: banking health telecom)")
    args = ap.parse_args()
    domains = set(args.domains)

    sources = collect_sources(domains)
    if not sources:
        print("No matching calls found. Check RESULTS_ROOT path.")
        return

    counts = {d: 0 for d in domains}
    total_sentences = 0

    for domain, src in sources:
        out_dir = OUTPUT_ROOT / domain
        out_dir.mkdir(parents=True, exist_ok=True)
        out_path = out_dir / src.name

        result = segment_call(src)
        with open(out_path, "w", encoding="utf-8") as f:
            json.dump(result, f, indent=2)

        counts[domain] += 1
        total_sentences += result["total"]
        print(f"  [{domain:8}]  {src.stem:<40}  "
              f"{result['agent_sentences']:>3}A + {result['customer_sentences']:>3}C "
              f"= {result['total']:>4} sentences  ->  {out_path.relative_to(OUTPUT_ROOT.parent.parent)}")

    print()
    print("-" * 72)
    for d in sorted(counts):
        print(f"  {d:10}  {counts[d]} calls")
    print(f"  {'TOTAL':10}  {sum(counts.values())} calls  Β·  {total_sentences} sentences")
    print(f"\nOutput root: {OUTPUT_ROOT}")


if __name__ == "__main__":
    main()