models¶
topic_segmentation.models
Topic-segmentation models of Yu et al. (2023) on BERT, RoBERTa, Longformer, and ModernBERT.
Attribute names (bert, roberta, longformer, model, dropout, loss_calculator)
match the state-dict keys of the trained checkpoints.
Segmentation
¶
Shared forward pass for backbones implementing encode().
Source code in src/topic_segmentation/models.py
14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 | |
heads(dropout)
¶
Source code in src/topic_segmentation/models.py
17 18 19 20 | |
forward(input_ids, attention_mask=None, token_type_ids=None, labels=None, extract_eop_segment_ids=None, eop_index_for_aggregate_batch_eop_features=None, sent_pair_orders=None, sent_token_mask=None)
¶
Return loss and anchor logits from inputs ordered as [anchor, augmented].
Source code in src/topic_segmentation/models.py
23 24 25 26 27 28 29 30 31 32 33 | |
Bert
¶
Bases: Segmentation, BertPreTrainedModel
Source code in src/topic_segmentation/models.py
36 37 38 39 40 41 42 43 44 45 | |
bert = BertModel(config)
instance-attribute
¶
encode(input_ids, attention_mask, token_type_ids=None)
¶
Source code in src/topic_segmentation/models.py
44 45 | |
Roberta
¶
Bases: Segmentation, RobertaPreTrainedModel
Source code in src/topic_segmentation/models.py
48 49 50 51 52 53 54 55 56 57 | |
roberta = RobertaModel(config, add_pooling_layer=False)
instance-attribute
¶
encode(input_ids, attention_mask, token_type_ids=None)
¶
Source code in src/topic_segmentation/models.py
56 57 | |
Longformer
¶
Bases: Segmentation, LongformerPreTrainedModel
Source code in src/topic_segmentation/models.py
60 61 62 63 64 65 66 67 68 69 70 | |
longformer = LongformerModel(config, add_pooling_layer=False)
instance-attribute
¶
encode(input_ids, attention_mask, token_type_ids=None)
¶
Source code in src/topic_segmentation/models.py
66 67 68 69 70 | |
ModernBert
¶
Bases: Segmentation, ModernBertPreTrainedModel
Source code in src/topic_segmentation/models.py
73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 | |
model = ModernBertModel(config)
instance-attribute
¶
encode(input_ids, attention_mask, token_type_ids=None)
¶
Source code in src/topic_segmentation/models.py
88 89 | |