Skip to content

augmentation

topic_segmentation.augmentation

Document augmentation (DA) and topic-aware sentence structure (TSSP) labels for the paired views.

Unit dataclass

Source code in src/topic_segmentation/augmentation.py
5
6
7
8
9
@dataclass
class Unit:
    text: str
    token_ids: list[int]
    label: int

text: str instance-attribute

token_ids: list[int] instance-attribute

label: int instance-attribute

AugmentedDocument dataclass

Source code in src/topic_segmentation/augmentation.py
12
13
14
15
16
17
@dataclass
class AugmentedDocument:
    input_ids: list[int]
    sentences: list[str]
    token_labels: list[int]
    pair_orders: list[int]

input_ids: list[int] instance-attribute

sentences: list[str] instance-attribute

token_labels: list[int] instance-attribute

pair_orders: list[int] instance-attribute

split_units(texts, labels, input_ids, boundary_ids)

Source code in src/topic_segmentation/augmentation.py
20
21
22
23
24
25
26
27
def split_units(texts, labels, input_ids, boundary_ids):
    starts = [i for i, token in enumerate(input_ids) if token in boundary_ids]
    assert len(starts) == len(texts)
    ends = starts[1:] + [len(input_ids)]
    return [
        Unit(text, input_ids[start:end], label)
        for text, label, start, end in zip(texts, labels, starts, ends)
    ]

section_spans(document, boundary_label)

Source code in src/topic_segmentation/augmentation.py
30
31
32
33
def section_spans(document, boundary_label):
    ends = [i for i, sentence in enumerate(document) if sentence.label == boundary_label]
    starts = [0] + [end + 1 for end in ends[:-1]]
    return list(zip(starts, ends))

mix_sections(documents, topic_spans, document_index, rng)

Shuffle topics, optionally drawing replacements from another document.

Source code in src/topic_segmentation/augmentation.py
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
def mix_sections(documents, topic_spans, document_index, rng):
    """Shuffle topics, optionally drawing replacements from another document."""
    topic_order = list(range(len(topic_spans[document_index])))
    rng.shuffle(topic_order)
    # Single-document batches also consume this draw before later shuffles.
    replace_topics = rng.random() > 0.5 and len(documents) > 1
    mixed = []
    for topic_index in topic_order:
        source_index = document_index
        prefix = f"{topic_index}-"
        if replace_topics and rng.random() > 0.5:
            candidates = [i for i in range(len(documents)) if i != document_index]
            source_index = rng.choice(candidates)
            topic_index = rng.choice(list(range(len(topic_spans[source_index]))))
            prefix = f"{source_index}-{topic_index}-"
        start, end = topic_spans[source_index][topic_index]
        for sentence_index in range(start, end + 1):
            sentence = documents[source_index][sentence_index]
            mixed.append(replace(sentence, text=f"{prefix}{sentence_index}-{sentence.text}"))
    return mixed

pair_order(position, sentence_indices)

TSSP label of a shuffled sentence: 2 opens a topic, 0 follows its predecessor, 1 otherwise.

Source code in src/topic_segmentation/augmentation.py
58
59
60
61
62
def pair_order(position, sentence_indices):
    """TSSP label of a shuffled sentence: 2 opens a topic, 0 follows its predecessor, 1 otherwise."""
    if position == 0:
        return 2
    return 0 if sentence_indices[position - 1] == sentence_indices[position] - 1 else 1

shuffle_units(document, label_to_id, rng)

Source code in src/topic_segmentation/augmentation.py
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
def shuffle_units(document, label_to_id, rng):
    output = AugmentedDocument(input_ids=[], sentences=[], token_labels=[], pair_orders=[])
    for start, end in section_spans(document, label_to_id["B-EOP"]):
        indices = list(range(start, end))
        rng.shuffle(indices)
        indices.append(end)  # The topic-final sentence always remains last.
        for position, sentence_index in enumerate(indices):
            sentence = document[sentence_index]
            label = label_to_id["B-EOP"] if position == len(indices) - 1 else label_to_id["O"]
            pair_label = pair_order(position, indices)
            ignored = [-100] * (len(sentence.token_ids) - 1)
            output.input_ids.extend(sentence.token_ids)
            output.sentences.append(f"{sentence_index}-{sentence.text}")
            output.token_labels.extend([label] + ignored)
            output.pair_orders.extend([pair_label] + ignored)
    return output

augment(sentences, labels, tokenized, label_to_id, boundary_ids, rng)

Create one augmented view per document, drawing from rng in document order.

Source code in src/topic_segmentation/augmentation.py
83
84
85
86
87
88
89
90
91
def augment(sentences, labels, tokenized, label_to_id, boundary_ids, rng):
    """Create one augmented view per document, drawing from `rng` in document order."""
    documents = [
        split_units(texts, [label_to_id.get(label, -100) for label in doc_labels], ids, boundary_ids)
        for texts, doc_labels, ids in zip(sentences, labels, tokenized["input_ids"])
    ]
    topic_spans = [section_spans(document, label_to_id["B-EOP"]) for document in documents]
    return [shuffle_units(mix_sections(documents, topic_spans, index, rng), label_to_id, rng)
            for index in range(len(documents))]