ManuelaDanu commited on
Commit
7598f04
·
verified ·
1 Parent(s): eb3ae38

Upload modeling.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. modeling.py +1778 -0
modeling.py ADDED
@@ -0,0 +1,1778 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from itertools import islice
2
+ from typing import Dict, List, Optional, Tuple, Union
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+ from torchcrf import CRF
7
+ from transformers import PretrainedConfig, PreTrainedModel
8
+ from transformers.modeling_outputs import TokenClassifierOutput
9
+
10
+ # Large negative number for masking impossible transitions
11
+ LARGE_NEGATIVE_NUMBER = -1e9
12
+ NUM_PER_LAYER = 16
13
+
14
+
15
+ class MultiHeadCRFConfig(PretrainedConfig):
16
+ """
17
+ Configuration class for Multi-Head CRF models.
18
+
19
+ Args:
20
+ entity_types: List of entity type names (e.g., ["DRUG", "DISEASE", "SYMPTOM"])
21
+ number_of_layers_per_head: Number of dense layers per head before classification
22
+ crf_reduction: Reduction mode for CRF loss ("mean", "sum", "token_mean", "none")
23
+ freeze_backbone: Whether to freeze the transformer backbone
24
+ num_frozen_encoders: Number of encoder layers to freeze (from bottom)
25
+ classifier_dropout: Dropout rate for classifier heads
26
+ **kwargs: Additional arguments passed to PretrainedConfig
27
+ """
28
+
29
+ model_type = "multihead-crf-tagger"
30
+
31
+ def __init__(
32
+ self,
33
+ entity_types: Optional[List[str]] = None,
34
+ number_of_layers_per_head: int = 1,
35
+ crf_reduction: str = "mean",
36
+ freeze_backbone: bool = False,
37
+ num_frozen_encoders: int = 0,
38
+ classifier_dropout: float = 0.1,
39
+ classifier_hidden_layers: Optional[Tuple] = None,
40
+ class_weights: Optional[List[float]] = None,
41
+ backbone_model_name: Optional[str] = None,
42
+ **kwargs,
43
+ ):
44
+ self.entity_types = entity_types or []
45
+ self.number_of_layers_per_head = number_of_layers_per_head
46
+ self.crf_reduction = crf_reduction
47
+ self.freeze_backbone = freeze_backbone
48
+ self.num_frozen_encoders = num_frozen_encoders
49
+ self.classifier_dropout = classifier_dropout
50
+ self.classifier_hidden_layers = classifier_hidden_layers
51
+ self.class_weights = class_weights
52
+ self.backbone_model_name = backbone_model_name
53
+ super().__init__(**kwargs)
54
+
55
+
56
+ class MultiHeadCRF(nn.Module):
57
+ """
58
+ Custom CRF implementation with BIO transition masking.
59
+
60
+ This CRF implementation includes:
61
+ - Proper initialization of transition parameters
62
+ - Masking of impossible BIO transitions (e.g., O -> I is invalid)
63
+ - Viterbi decoding for inference
64
+
65
+ Args:
66
+ num_tags: Number of tags (typically 3 for BIO: O, B, I)
67
+ batch_first: Whether batch dimension is first
68
+ """
69
+
70
+ def __init__(self, num_tags: int, batch_first: bool = True) -> None:
71
+ if num_tags <= 0:
72
+ raise ValueError(f"invalid number of tags: {num_tags}")
73
+ super().__init__()
74
+ self.num_tags = num_tags
75
+ self.batch_first = batch_first
76
+ self.start_transitions = nn.Parameter(torch.empty(num_tags))
77
+ self.end_transitions = nn.Parameter(torch.empty(num_tags))
78
+ self.transitions = nn.Parameter(torch.empty(num_tags, num_tags))
79
+
80
+ self.reset_parameters()
81
+ self.mask_impossible_transitions()
82
+
83
+ def reset_parameters(self) -> None:
84
+ """Initialize the transition parameters uniformly between -0.1 and 0.1."""
85
+ nn.init.uniform_(self.start_transitions, -0.1, 0.1)
86
+ nn.init.uniform_(self.end_transitions, -0.1, 0.1)
87
+ nn.init.uniform_(self.transitions, -0.1, 0.1)
88
+
89
+ def mask_impossible_transitions(self) -> None:
90
+ """
91
+ Set impossible BIO transitions to large negative values.
92
+
93
+ For standard BIO tagging with tags [O=0, B=1, I=2]:
94
+ - Cannot start with I tag
95
+ - Cannot transition from O to I
96
+ """
97
+ with torch.no_grad():
98
+ # Assuming BIO scheme: O=0, B=1, I=2
99
+ # Cannot start with I
100
+ if self.num_tags > 2:
101
+ self.start_transitions[2] = LARGE_NEGATIVE_NUMBER
102
+ # Cannot go from O to I
103
+ self.transitions[0][2] = LARGE_NEGATIVE_NUMBER
104
+
105
+ # If PADDING token exists (index 3+), mask its transitions
106
+ if self.num_tags > 3:
107
+ # Cannot start with PADDING
108
+ self.start_transitions[3] = LARGE_NEGATIVE_NUMBER
109
+ # Cannot transition to PADDING from valid tags
110
+ for i in range(3):
111
+ self.transitions[i][3] = LARGE_NEGATIVE_NUMBER
112
+ # Cannot transition from PADDING to valid tags
113
+ for i in range(3):
114
+ self.transitions[3][i] = LARGE_NEGATIVE_NUMBER
115
+
116
+ def __repr__(self) -> str:
117
+ return f"{self.__class__.__name__}(num_tags={self.num_tags})"
118
+
119
+ def forward(
120
+ self,
121
+ emissions: torch.Tensor,
122
+ tags: torch.Tensor,
123
+ mask: Optional[torch.Tensor] = None,
124
+ reduction: str = "mean",
125
+ ) -> torch.Tensor:
126
+ """
127
+ Compute the negative log likelihood of a sequence of tags given emission scores.
128
+
129
+ Args:
130
+ emissions: Emission scores (batch_size, seq_length, num_tags) if batch_first
131
+ tags: Gold tag sequence (batch_size, seq_length) if batch_first
132
+ mask: Mask tensor (batch_size, seq_length) if batch_first
133
+ reduction: Loss reduction mode ("none", "sum", "mean", "token_mean")
134
+
135
+ Returns:
136
+ Negative log likelihood loss
137
+ """
138
+ self._validate(emissions, tags=tags, mask=mask)
139
+ if reduction not in ("none", "sum", "mean", "token_mean"):
140
+ raise ValueError(f"invalid reduction: {reduction}")
141
+ if mask is None:
142
+ mask = torch.ones_like(tags, dtype=torch.uint8)
143
+
144
+ # Ensure all tensors are on the same device as emissions
145
+ device = emissions.device
146
+ tags = tags.to(device)
147
+ mask = mask.to(device)
148
+
149
+ if self.batch_first:
150
+ emissions = emissions.transpose(0, 1)
151
+ tags = tags.transpose(0, 1)
152
+ mask = mask.transpose(0, 1)
153
+
154
+ # shape: (batch_size,)
155
+ numerator = self._compute_score(emissions, tags, mask)
156
+ # shape: (batch_size,)
157
+ denominator = self._compute_normalizer(emissions, mask)
158
+ # shape: (batch_size,)
159
+ llh = numerator - denominator
160
+ nllh = -llh
161
+
162
+ if reduction == "none":
163
+ return nllh
164
+ if reduction == "sum":
165
+ return nllh.sum()
166
+ if reduction == "mean":
167
+ return nllh.mean()
168
+ assert reduction == "token_mean"
169
+ return nllh.sum() / mask.type_as(emissions).sum()
170
+
171
+ def decode(
172
+ self, emissions: torch.Tensor, mask: Optional[torch.Tensor] = None
173
+ ) -> List[List[int]]:
174
+ """
175
+ Find the most likely tag sequence using Viterbi algorithm.
176
+
177
+ Args:
178
+ emissions: Emission scores
179
+ mask: Mask tensor
180
+
181
+ Returns:
182
+ List of best tag sequences for each batch
183
+ """
184
+ self._validate(emissions, mask=mask)
185
+ if mask is None:
186
+ mask = emissions.new_ones(emissions.shape[:2], dtype=torch.uint8)
187
+
188
+ if self.batch_first:
189
+ emissions = emissions.transpose(0, 1)
190
+ mask = mask.transpose(0, 1)
191
+
192
+ return self._viterbi_decode(emissions, mask)
193
+
194
+ def _validate(
195
+ self,
196
+ emissions: torch.Tensor,
197
+ tags: Optional[torch.Tensor] = None,
198
+ mask: Optional[torch.Tensor] = None,
199
+ ) -> None:
200
+ if emissions.dim() != 3:
201
+ raise ValueError(
202
+ f"emissions must have dimension of 3, got {emissions.dim()}"
203
+ )
204
+ if emissions.size(2) != self.num_tags:
205
+ raise ValueError(
206
+ f"expected last dimension of emissions is {self.num_tags}, "
207
+ f"got {emissions.size(2)}"
208
+ )
209
+
210
+ if tags is not None:
211
+ if emissions.shape[:2] != tags.shape:
212
+ raise ValueError(
213
+ "the first two dimensions of emissions and tags must match, "
214
+ f"got {tuple(emissions.shape[:2])} and {tuple(tags.shape)}"
215
+ )
216
+
217
+ if mask is not None:
218
+ if emissions.shape[:2] != mask.shape:
219
+ raise ValueError(
220
+ "the first two dimensions of emissions and mask must match, "
221
+ f"got {tuple(emissions.shape[:2])} and {tuple(mask.shape)}"
222
+ )
223
+ no_empty_seq = not self.batch_first and mask[0].all()
224
+ no_empty_seq_bf = self.batch_first and mask[:, 0].all()
225
+ if not no_empty_seq and not no_empty_seq_bf:
226
+ raise ValueError("mask of the first timestep must all be on")
227
+
228
+ def _compute_score(
229
+ self, emissions: torch.Tensor, tags: torch.Tensor, mask: torch.Tensor
230
+ ) -> torch.Tensor:
231
+ # emissions: (seq_length, batch_size, num_tags)
232
+ # tags: (seq_length, batch_size)
233
+ # mask: (seq_length, batch_size)
234
+ assert emissions.dim() == 3 and tags.dim() == 2
235
+ assert emissions.shape[:2] == tags.shape
236
+ assert emissions.size(2) == self.num_tags
237
+ assert mask.shape == tags.shape
238
+ assert mask[0].all()
239
+
240
+ # Move all tensors to the same device as emissions
241
+ device = emissions.device
242
+ tags = tags.to(device)
243
+ mask = mask.to(device)
244
+
245
+ seq_length, batch_size = tags.shape
246
+ mask = mask.type_as(emissions)
247
+
248
+ # Start transition score and first emission
249
+ # Ensure arange is on the same device as other tensors
250
+ batch_indices = torch.arange(batch_size, device=device)
251
+ score = self.start_transitions[tags[0]]
252
+ score += emissions[0, batch_indices, tags[0]]
253
+
254
+ for i in range(1, seq_length):
255
+ score += self.transitions[tags[i - 1], tags[i]] * mask[i]
256
+ score += emissions[i, batch_indices, tags[i]] * mask[i]
257
+
258
+ # End transition score
259
+ seq_ends = mask.long().sum(dim=0) - 1
260
+ last_tags = tags[seq_ends, batch_indices]
261
+ score += self.end_transitions[last_tags]
262
+
263
+ return score
264
+
265
+ def _compute_normalizer(
266
+ self, emissions: torch.Tensor, mask: torch.Tensor
267
+ ) -> torch.Tensor:
268
+ # emissions: (seq_length, batch_size, num_tags)
269
+ # mask: (seq_length, batch_size)
270
+ assert emissions.dim() == 3 and mask.dim() == 2
271
+ assert emissions.shape[:2] == mask.shape
272
+ assert emissions.size(2) == self.num_tags
273
+ assert mask[0].all()
274
+
275
+ seq_length = emissions.size(0)
276
+
277
+ # Start transition score and first emission
278
+ score = self.start_transitions + emissions[0]
279
+
280
+ for i in range(1, seq_length):
281
+ broadcast_score = score.unsqueeze(2)
282
+ broadcast_emissions = emissions[i].unsqueeze(1)
283
+ next_score = broadcast_score + self.transitions + broadcast_emissions
284
+ next_score = torch.logsumexp(next_score, dim=1)
285
+ score = torch.where(mask[i].unsqueeze(1).bool(), next_score, score)
286
+
287
+ score += self.end_transitions
288
+ return torch.logsumexp(score, dim=1)
289
+
290
+ def _viterbi_decode(
291
+ self, emissions: torch.Tensor, mask: torch.Tensor
292
+ ) -> List[List[int]]:
293
+ # emissions: (seq_length, batch_size, num_tags)
294
+ # mask: (seq_length, batch_size)
295
+ assert emissions.dim() == 3 and mask.dim() == 2
296
+ assert emissions.shape[:2] == mask.shape
297
+ assert emissions.size(2) == self.num_tags
298
+ assert mask[0].all()
299
+
300
+ seq_length, batch_size = mask.shape
301
+
302
+ # Start transition and first emission
303
+ score = self.start_transitions + emissions[0]
304
+ history = []
305
+
306
+ for i in range(1, seq_length):
307
+ broadcast_score = score.unsqueeze(2)
308
+ broadcast_emission = emissions[i].unsqueeze(1)
309
+ next_score = broadcast_score + self.transitions + broadcast_emission
310
+ next_score, indices = next_score.max(dim=1)
311
+ score = torch.where(mask[i].unsqueeze(1).bool(), next_score, score)
312
+ history.append(indices)
313
+
314
+ score += self.end_transitions
315
+
316
+ # Trace back
317
+ seq_ends = mask.long().sum(dim=0) - 1
318
+ best_tags_list = []
319
+
320
+ for idx in range(batch_size):
321
+ _, best_last_tag = score[idx].max(dim=0)
322
+ best_tags = [best_last_tag.item()]
323
+
324
+ for hist in reversed(history[: seq_ends[idx]]):
325
+ best_last_tag = hist[idx][best_tags[-1]]
326
+ best_tags.append(best_last_tag.item())
327
+
328
+ best_tags.reverse()
329
+ best_tags_list.append(best_tags)
330
+
331
+ return best_tags_list
332
+
333
+
334
+ class TokenClassificationModelCRF(PreTrainedModel):
335
+ """
336
+ Custom token classification model with CRF layer and configurable classifier head.
337
+ This model can be loaded with trust_remote_code=True for HuggingFace Hub compatibility.
338
+ """
339
+
340
+ def __init__(
341
+ self,
342
+ config,
343
+ base_model=None,
344
+ freeze_backbone=False,
345
+ classifier_hidden_layers=None,
346
+ classifier_dropout=0.1,
347
+ ):
348
+ super().__init__(config)
349
+ self.config = config
350
+ self.num_labels = config.num_labels
351
+
352
+ # If base_model is not provided, load it from config
353
+ if base_model is None:
354
+ from transformers import AutoConfig, RobertaForTokenClassification
355
+
356
+ # Use backbone_model_name if available, fallback to name_or_path
357
+ # This is critical because name_or_path gets overwritten during save/load
358
+ backbone_name = getattr(config, "backbone_model_name", None)
359
+ if backbone_name is None:
360
+ backbone_name = getattr(config, "name_or_path", None) or getattr(
361
+ config, "_name_or_path", None
362
+ )
363
+ if backbone_name is None:
364
+ raise ValueError(
365
+ "config.backbone_model_name (or config.name_or_path) is required to load pretrained backbone"
366
+ )
367
+
368
+ # Create a clean config for the backbone
369
+ backbone_config = AutoConfig.from_pretrained(backbone_name)
370
+ backbone_config.hidden_dropout_prob = getattr(
371
+ config, "hidden_dropout_prob", 0.1
372
+ )
373
+ backbone_config.num_labels = config.num_labels
374
+
375
+ roberta_model = RobertaForTokenClassification.from_pretrained(
376
+ backbone_name, config=backbone_config
377
+ )
378
+ self.roberta = roberta_model.roberta
379
+
380
+ # Store backbone_model_name in config for future loading
381
+ if (
382
+ not hasattr(config, "backbone_model_name")
383
+ or config.backbone_model_name is None
384
+ ):
385
+ config.backbone_model_name = backbone_name
386
+ else:
387
+ if hasattr(base_model, "roberta"):
388
+ self.roberta = base_model.roberta
389
+ else:
390
+ self.roberta = base_model
391
+
392
+ self.lm_output_size = self.roberta.config.hidden_size
393
+
394
+ # Store configuration for saving/loading
395
+ self.config.freeze_backbone = freeze_backbone
396
+ self.config.classifier_hidden_layers = classifier_hidden_layers
397
+ self.config.classifier_dropout = classifier_dropout
398
+
399
+ if freeze_backbone:
400
+ print("+" * 30, "\n\n", "Freezing backbone...", "+" * 30, "\n\n")
401
+ for param in self.roberta.parameters():
402
+ param.requires_grad = False
403
+ self.roberta.eval()
404
+ else:
405
+ print("+" * 30, "\n\n", "NOT Freezing backbone...", "+" * 30, "\n\n")
406
+
407
+ self.roberta.train(not freeze_backbone)
408
+
409
+ self.dropout = nn.Dropout(
410
+ config.hidden_dropout_prob
411
+ if hasattr(config, "hidden_dropout_prob")
412
+ else 0.1
413
+ )
414
+ self.crf = CRF(self.num_labels, batch_first=True)
415
+
416
+ self._build_classifier_head(classifier_hidden_layers, classifier_dropout)
417
+
418
+ def _build_classifier_head(self, hidden_layers, dropout_rate):
419
+ """
420
+ Build a flexible classifier head with configurable hidden layers and dropout.
421
+
422
+ Args:
423
+ hidden_layers: Tuple of integers representing the number of neurons in each hidden layer.
424
+ None or empty tuple means a simple linear layer.
425
+ dropout_rate: Dropout probability between layers
426
+ """
427
+ layers = []
428
+ input_size = self.lm_output_size
429
+
430
+ # If hidden_layers is None or empty, just create a simple linear layer
431
+ if not hidden_layers:
432
+ self.classifier = nn.Sequential(
433
+ nn.Dropout(dropout_rate), nn.Linear(input_size, self.num_labels)
434
+ )
435
+ return
436
+
437
+ # Build MLP with specified hidden layers
438
+ for hidden_size in hidden_layers:
439
+ layers.append(nn.Linear(input_size, hidden_size))
440
+ layers.append(nn.ReLU())
441
+ layers.append(nn.Dropout(dropout_rate))
442
+ input_size = hidden_size
443
+
444
+ # Final classification layer
445
+ layers.append(nn.Linear(input_size, self.num_labels))
446
+
447
+ # Create sequential model
448
+ self.classifier = nn.Sequential(*layers)
449
+
450
+ def forward(
451
+ self,
452
+ input_ids: Optional[torch.LongTensor] = None,
453
+ attention_mask: Optional[torch.FloatTensor] = None,
454
+ token_type_ids: Optional[torch.LongTensor] = None,
455
+ position_ids: Optional[torch.LongTensor] = None,
456
+ head_mask: Optional[torch.FloatTensor] = None,
457
+ inputs_embeds: Optional[torch.FloatTensor] = None,
458
+ labels: Optional[torch.LongTensor] = None,
459
+ output_attentions: Optional[bool] = None,
460
+ output_hidden_states: Optional[bool] = None,
461
+ return_dict: Optional[bool] = None,
462
+ **kwargs,
463
+ ) -> Union[Tuple[torch.Tensor], TokenClassifierOutput]:
464
+ return_dict = (
465
+ return_dict if return_dict is not None else self.config.use_return_dict
466
+ )
467
+
468
+ outputs = self.roberta(
469
+ input_ids,
470
+ attention_mask=attention_mask,
471
+ token_type_ids=token_type_ids,
472
+ position_ids=position_ids,
473
+ head_mask=head_mask,
474
+ inputs_embeds=inputs_embeds,
475
+ output_attentions=output_attentions,
476
+ output_hidden_states=output_hidden_states,
477
+ return_dict=return_dict,
478
+ )
479
+
480
+ sequence_output = self.dropout(outputs.last_hidden_state)
481
+ logits = self.classifier(sequence_output) # Emissions for CRF
482
+
483
+ loss = None
484
+ if labels is not None:
485
+ # CRF calculates the log-likelihood of the correct sequence
486
+ # We use a negative sign to convert it into a loss
487
+ # Following ieeta-pt approach: don't pass mask to CRF
488
+ # All positions have valid labels (O for special/padding tokens)
489
+ labels_long = labels.long()
490
+ loss = -self.crf(logits, labels_long, reduction="mean")
491
+
492
+ if not return_dict:
493
+ output = (logits,) + outputs[2:]
494
+ return ((loss,) + output) if loss is not None else output
495
+
496
+ return TokenClassifierOutput(
497
+ loss=loss,
498
+ logits=logits,
499
+ hidden_states=outputs.hidden_states,
500
+ attentions=outputs.attentions,
501
+ )
502
+
503
+ @property
504
+ def device_info(self):
505
+ return next(self.parameters()).device
506
+
507
+ def get_input_embeddings(self):
508
+ return self.roberta.get_input_embeddings()
509
+
510
+ def set_input_embeddings(self, value):
511
+ self.roberta.set_input_embeddings(value)
512
+
513
+ @classmethod
514
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
515
+ """Override from_pretrained to handle custom model loading"""
516
+ config = kwargs.pop("config", None)
517
+ if config is None:
518
+ from transformers import AutoConfig
519
+
520
+ config = AutoConfig.from_pretrained(pretrained_model_name_or_path, **kwargs)
521
+
522
+ # Extract custom parameters from config if they exist
523
+ freeze_backbone = getattr(config, "freeze_backbone", False)
524
+ classifier_hidden_layers = getattr(config, "classifier_hidden_layers", None)
525
+ classifier_dropout = getattr(config, "classifier_dropout", 0.1)
526
+
527
+ model = cls(
528
+ config=config,
529
+ freeze_backbone=freeze_backbone,
530
+ classifier_hidden_layers=classifier_hidden_layers,
531
+ classifier_dropout=classifier_dropout,
532
+ )
533
+
534
+ # Load state dict if available
535
+ try:
536
+ state_dict = torch.load(
537
+ f"{pretrained_model_name_or_path}/pytorch_model.bin", map_location="cpu"
538
+ )
539
+ model.load_state_dict(state_dict)
540
+ except:
541
+ # If loading fails, the model will be initialized with random weights
542
+ print(
543
+ "Warning: Could not load pre-trained weights. Using randomly initialized model."
544
+ )
545
+
546
+ return model
547
+
548
+
549
+ class TokenClassificationModelMultiHeadCRF(PreTrainedModel):
550
+ """
551
+ Multi-Head CRF model for token classification with multiple entity types.
552
+
553
+ Each entity type gets its own classification head and CRF layer, allowing
554
+ for independent BIO tagging per entity type. This is useful for scenarios
555
+ where entities can overlap or when different entity types have different
556
+ transition patterns.
557
+
558
+ Args:
559
+ config: MultiHeadCRFConfig or compatible config with entity_types
560
+ base_model: Optional pre-trained RoBERTa model
561
+ freeze_backbone: Whether to freeze transformer weights
562
+ """
563
+
564
+ config_class = MultiHeadCRFConfig
565
+ base_model_prefix = "roberta"
566
+ _keys_to_ignore_on_load_unexpected = [r"pooler"]
567
+
568
+ def __init__(self, config, base_model=None, freeze_backbone=None):
569
+ super().__init__(config)
570
+ self.config = config
571
+
572
+ # Get entity types from config
573
+ self.entity_types = getattr(config, "entity_types", [])
574
+ if not self.entity_types:
575
+ raise ValueError("entity_types must be provided in config")
576
+
577
+ # Number of labels per head (typically 3 for BIO: O, B, I) + padding
578
+ self.num_labels = config.num_labels
579
+ # self.num_labels_with_pad = self.num_labels + 1 # should be self.num_labels?
580
+
581
+ # Configuration parameters
582
+ self.number_of_layers_per_head = getattr(config, "number_of_layers_per_head", 1)
583
+ self.crf_reduction = getattr(config, "crf_reduction", "mean")
584
+ freeze_backbone = (
585
+ freeze_backbone
586
+ if freeze_backbone is not None
587
+ else getattr(config, "freeze_backbone", False)
588
+ )
589
+ self.num_frozen_encoders = getattr(config, "num_frozen_encoders", 0)
590
+ classifier_dropout = getattr(config, "classifier_dropout", 0.1)
591
+
592
+ # Initialize the transformer backbone
593
+ if base_model is None:
594
+ from transformers import AutoConfig, RobertaModel
595
+
596
+ # Use backbone_model_name if available, fallback to name_or_path
597
+ backbone_name = getattr(config, "backbone_model_name", None)
598
+ if backbone_name is None:
599
+ backbone_name = getattr(config, "name_or_path", None) or getattr(
600
+ config, "_name_or_path", None
601
+ )
602
+
603
+ if backbone_name:
604
+ # Load pretrained weights
605
+ backbone_config = AutoConfig.from_pretrained(backbone_name)
606
+ backbone_config.hidden_dropout_prob = getattr(
607
+ config, "hidden_dropout_prob", 0.1
608
+ )
609
+ self.roberta = RobertaModel.from_pretrained(
610
+ backbone_name, config=backbone_config, add_pooling_layer=False
611
+ )
612
+ # Store backbone_model_name in config for future loading
613
+ if (
614
+ not hasattr(config, "backbone_model_name")
615
+ or config.backbone_model_name is None
616
+ ):
617
+ config.backbone_model_name = backbone_name
618
+ else:
619
+ # Fallback: initialize without pretrained weights (not recommended)
620
+ self.roberta = RobertaModel(config, add_pooling_layer=False)
621
+ else:
622
+ if hasattr(base_model, "roberta"):
623
+ self.roberta = base_model.roberta
624
+ else:
625
+ self.roberta = base_model
626
+
627
+ self.hidden_size = config.hidden_size
628
+ self.dropout = nn.Dropout(getattr(config, "hidden_dropout_prob", 0.1))
629
+
630
+ # Create heads for each entity type
631
+ print(f"Creating Multi-Head CRF with entity types: {sorted(self.entity_types)}")
632
+
633
+ for entity_type in self.entity_types:
634
+ # Dense layers per head
635
+ for i in range(self.number_of_layers_per_head):
636
+ setattr(
637
+ self,
638
+ f"{entity_type}_dense_{i}",
639
+ nn.Linear(self.hidden_size, self.hidden_size),
640
+ )
641
+ setattr(
642
+ self,
643
+ f"{entity_type}_dense_activation_{i}",
644
+ nn.GELU(approximate="none"),
645
+ )
646
+ setattr(
647
+ self, f"{entity_type}_dropout_{i}", nn.Dropout(classifier_dropout)
648
+ )
649
+
650
+ # Classifier and CRF per head
651
+ setattr(
652
+ self,
653
+ f"{entity_type}_classifier",
654
+ nn.Linear(self.hidden_size, self.num_labels),
655
+ )
656
+ setattr(
657
+ self,
658
+ f"{entity_type}_crf",
659
+ MultiHeadCRF(num_tags=self.num_labels, batch_first=True),
660
+ )
661
+
662
+ # Handle freezing
663
+ if freeze_backbone:
664
+ self._freeze_backbone()
665
+
666
+ def _freeze_backbone(self):
667
+ """Freeze transformer backbone parameters."""
668
+ print("+" * 30, "\n\n", "Freezing backbone...", "+" * 30, "\n\n")
669
+
670
+ # Freeze embeddings
671
+ for param in self.roberta.embeddings.parameters():
672
+ param.requires_grad = False
673
+
674
+ # Optionally freeze some encoder layers
675
+ if self.num_frozen_encoders > 0:
676
+ for _, param in islice(
677
+ self.roberta.encoder.named_parameters(),
678
+ self.num_frozen_encoders * NUM_PER_LAYER,
679
+ ):
680
+ param.requires_grad = False
681
+
682
+ def reset_head_parameters(self):
683
+ """Reset parameters for all heads (useful after loading pretrained weights)."""
684
+ for entity_type in self.entity_types:
685
+ for i in range(self.number_of_layers_per_head):
686
+ getattr(self, f"{entity_type}_dense_{i}").reset_parameters()
687
+ getattr(self, f"{entity_type}_classifier").reset_parameters()
688
+ getattr(self, f"{entity_type}_crf").reset_parameters()
689
+ getattr(self, f"{entity_type}_crf").mask_impossible_transitions()
690
+
691
+ def forward(
692
+ self,
693
+ input_ids: Optional[torch.LongTensor] = None,
694
+ attention_mask: Optional[torch.FloatTensor] = None,
695
+ token_type_ids: Optional[torch.LongTensor] = None,
696
+ position_ids: Optional[torch.LongTensor] = None,
697
+ head_mask: Optional[torch.FloatTensor] = None,
698
+ inputs_embeds: Optional[torch.FloatTensor] = None,
699
+ labels: Optional[Dict[str, torch.LongTensor]] = None,
700
+ output_attentions: Optional[bool] = None,
701
+ output_hidden_states: Optional[bool] = None,
702
+ return_dict: Optional[bool] = None,
703
+ **kwargs,
704
+ ):
705
+ """
706
+ Forward pass through the multi-head CRF model.
707
+
708
+ Args:
709
+ input_ids: Input token IDs
710
+ attention_mask: Attention mask
711
+ labels: Dictionary mapping entity types to label tensors
712
+ e.g., {"DRUG": tensor, "DISEASE": tensor}
713
+ ... other standard transformer arguments
714
+
715
+ Returns:
716
+ During training (labels provided):
717
+ Tuple of (total_loss, logits_dict) where logits_dict maps entity types to logits
718
+ During inference (no labels):
719
+ List of prediction tensors, one per entity type (sorted alphabetically)
720
+ """
721
+ return_dict = (
722
+ return_dict if return_dict is not None else self.config.use_return_dict
723
+ )
724
+
725
+ # Get transformer outputs
726
+ outputs = self.roberta(
727
+ input_ids,
728
+ attention_mask=attention_mask,
729
+ token_type_ids=token_type_ids,
730
+ position_ids=position_ids,
731
+ head_mask=head_mask,
732
+ inputs_embeds=inputs_embeds,
733
+ output_attentions=output_attentions,
734
+ output_hidden_states=output_hidden_states,
735
+ return_dict=return_dict,
736
+ )
737
+
738
+ sequence_output = outputs[0]
739
+ sequence_output = self.dropout(sequence_output) # (batch, seq_len, hidden)
740
+
741
+ # Compute logits for each head
742
+ logits = {}
743
+ for entity_type in self.entity_types:
744
+ head_output = sequence_output
745
+ for i in range(self.number_of_layers_per_head):
746
+ head_output = getattr(self, f"{entity_type}_dense_{i}")(head_output)
747
+ head_output = getattr(self, f"{entity_type}_dense_activation_{i}")(
748
+ head_output
749
+ )
750
+ head_output = getattr(self, f"{entity_type}_dropout_{i}")(head_output)
751
+ logits[entity_type] = getattr(self, f"{entity_type}_classifier")(
752
+ head_output
753
+ )
754
+
755
+ if labels is not None:
756
+ # Training mode - compute CRF loss for each head
757
+ # Following ieeta-pt approach: don't pass mask to CRF
758
+ # All positions have valid labels (O for special/padding tokens)
759
+ losses = {}
760
+
761
+ for entity_type in self.entity_types:
762
+ if entity_type in labels:
763
+ # Ensure labels are on the same device as logits
764
+ entity_labels = (
765
+ labels[entity_type].long().to(logits[entity_type].device)
766
+ )
767
+ crf = getattr(self, f"{entity_type}_crf")
768
+ # CRF returns negative log likelihood, we want to minimize it
769
+ losses[entity_type] = crf(
770
+ logits[entity_type],
771
+ entity_labels,
772
+ reduction=self.crf_reduction,
773
+ )
774
+
775
+ # Sum losses from all heads
776
+ total_loss = sum(losses.values())
777
+ return total_loss, logits
778
+
779
+ else:
780
+ # Inference mode - decode each head
781
+ # Following ieeta-pt approach: don't pass mask, decode all positions
782
+ predictions = {}
783
+
784
+ for entity_type in self.entity_types:
785
+ crf = getattr(self, f"{entity_type}_crf")
786
+ decoded = crf.decode(logits[entity_type])
787
+ predictions[entity_type] = torch.tensor(decoded)
788
+
789
+ # Return as list sorted by entity type for consistency
790
+ return [predictions[ent] for ent in sorted(self.entity_types)]
791
+
792
+ def get_input_embeddings(self):
793
+ return self.roberta.get_input_embeddings()
794
+
795
+ def set_input_embeddings(self, value):
796
+ self.roberta.set_input_embeddings(value)
797
+
798
+ @classmethod
799
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
800
+ """Override from_pretrained to handle custom model loading."""
801
+ import json
802
+ import os
803
+
804
+ config = kwargs.pop("config", None)
805
+
806
+ if config is None:
807
+ # Load config directly from JSON to get all saved attributes
808
+ config_file = os.path.join(pretrained_model_name_or_path, "config.json")
809
+ if os.path.exists(config_file):
810
+ with open(config_file, "r") as f:
811
+ config_dict = json.load(f)
812
+
813
+ # Create MultiHeadCRFConfig with all loaded parameters
814
+ config = MultiHeadCRFConfig(**config_dict)
815
+ else:
816
+ from transformers import AutoConfig
817
+
818
+ config = AutoConfig.from_pretrained(
819
+ pretrained_model_name_or_path,
820
+ trust_remote_code=kwargs.get("trust_remote_code", True),
821
+ )
822
+
823
+ # Ensure config has all required RoBERTa parameters
824
+ # These are needed to initialize RobertaModel
825
+ roberta_defaults = {
826
+ # Core model architecture
827
+ "layer_norm_eps": 1e-5,
828
+ "hidden_size": 768,
829
+ "num_hidden_layers": 12,
830
+ "num_attention_heads": 12,
831
+ "intermediate_size": 3072,
832
+ "hidden_act": "gelu",
833
+ "hidden_dropout_prob": 0.1,
834
+ "attention_probs_dropout_prob": 0.1,
835
+ "max_position_embeddings": 514,
836
+ "type_vocab_size": 1,
837
+ "initializer_range": 0.02,
838
+ "vocab_size": 52000,
839
+ # Token IDs
840
+ "pad_token_id": 1,
841
+ "bos_token_id": 0,
842
+ "eos_token_id": 2,
843
+ # Position embeddings
844
+ "position_embedding_type": "absolute",
845
+ # Model behavior flags
846
+ "use_cache": True,
847
+ "is_decoder": False,
848
+ "add_cross_attention": False,
849
+ "chunk_size_feed_forward": 0,
850
+ "output_hidden_states": False,
851
+ "output_attentions": False,
852
+ "torchscript": False,
853
+ "tie_word_embeddings": True,
854
+ "return_dict": True,
855
+ # Gradient checkpointing
856
+ "gradient_checkpointing": False,
857
+ # Pruning
858
+ "pruned_heads": {},
859
+ # Problem type (for classification)
860
+ "problem_type": None,
861
+ # Embedding layer norm
862
+ "embedding_size": None,
863
+ }
864
+
865
+ for key, default_value in roberta_defaults.items():
866
+ if not hasattr(config, key) or getattr(config, key) is None:
867
+ setattr(config, key, default_value)
868
+
869
+ freeze_backbone = getattr(config, "freeze_backbone", False)
870
+
871
+ model = cls(config=config, freeze_backbone=freeze_backbone)
872
+
873
+ # Load state dict if available
874
+ weight_file = os.path.join(pretrained_model_name_or_path, "pytorch_model.bin")
875
+ safetensors_file = os.path.join(
876
+ pretrained_model_name_or_path, "model.safetensors"
877
+ )
878
+
879
+ try:
880
+ if os.path.exists(safetensors_file):
881
+ from safetensors.torch import load_file
882
+
883
+ state_dict = load_file(safetensors_file)
884
+ model.load_state_dict(state_dict)
885
+ elif os.path.exists(weight_file):
886
+ state_dict = torch.load(weight_file, map_location="cpu")
887
+ model.load_state_dict(state_dict)
888
+ else:
889
+ print(
890
+ "Warning: No pre-trained weights found. Using randomly initialized model."
891
+ )
892
+ except Exception as e:
893
+ print(f"Warning: Could not load pre-trained weights: {e}")
894
+
895
+ return model
896
+
897
+
898
+ class MultiHeadConfig(PretrainedConfig):
899
+ """
900
+ Configuration class for Multi-Head models (without CRF).
901
+
902
+ Args:
903
+ entity_types: List of entity type names (e.g., ["DRUG", "DISEASE", "SYMPTOM"])
904
+ number_of_layers_per_head: Number of dense layers per head before classification
905
+ freeze_backbone: Whether to freeze the transformer backbone
906
+ num_frozen_encoders: Number of encoder layers to freeze (from bottom)
907
+ classifier_dropout: Dropout rate for classifier heads
908
+ use_class_weights: Whether to use class weights for loss computation
909
+ class_weights: Optional dict mapping entity types to weight lists
910
+ **kwargs: Additional arguments passed to PretrainedConfig
911
+ """
912
+
913
+ model_type = "multihead-tagger"
914
+
915
+ def __init__(
916
+ self,
917
+ entity_types: Optional[List[str]] = None,
918
+ number_of_layers_per_head: int = 1,
919
+ freeze_backbone: bool = False,
920
+ num_frozen_encoders: int = 0,
921
+ classifier_dropout: float = 0.1,
922
+ use_class_weights: bool = False,
923
+ class_weights: Optional[Dict[str, List[float]]] = None,
924
+ backbone_model_name: Optional[str] = None,
925
+ **kwargs,
926
+ ):
927
+ self.entity_types = entity_types or []
928
+ self.number_of_layers_per_head = number_of_layers_per_head
929
+ self.freeze_backbone = freeze_backbone
930
+ self.num_frozen_encoders = num_frozen_encoders
931
+ self.classifier_dropout = classifier_dropout
932
+ self.use_class_weights = use_class_weights
933
+ self.class_weights = class_weights
934
+ self.backbone_model_name = backbone_model_name
935
+ super().__init__(**kwargs)
936
+
937
+
938
+ class TokenClassificationModelMultiHead(PreTrainedModel):
939
+ """
940
+ Multi-Head model for token classification with multiple entity types (no CRF).
941
+
942
+ Each entity type gets its own classification head, allowing for independent
943
+ BIO tagging per entity type. This is useful for scenarios where entities
944
+ can overlap or when different entity types need separate classification.
945
+
946
+ Unlike the CRF variant, this model uses standard CrossEntropyLoss and
947
+ argmax decoding, which is faster but doesn't enforce valid BIO sequences.
948
+
949
+ Args:
950
+ config: MultiHeadConfig or compatible config with entity_types
951
+ base_model: Optional pre-trained RoBERTa model
952
+ freeze_backbone: Whether to freeze transformer weights
953
+ """
954
+
955
+ config_class = MultiHeadConfig
956
+ base_model_prefix = "roberta"
957
+ _keys_to_ignore_on_load_unexpected = [r"pooler"]
958
+
959
+ def __init__(self, config, base_model=None, freeze_backbone=None):
960
+ super().__init__(config)
961
+ self.config = config
962
+
963
+ # Get entity types from config
964
+ self.entity_types = getattr(config, "entity_types", [])
965
+ if not self.entity_types:
966
+ raise ValueError("entity_types must be provided in config")
967
+
968
+ # Number of labels per head (typically 3 for BIO: O, B, I)
969
+ self.num_labels = config.num_labels
970
+
971
+ # Configuration parameters
972
+ self.number_of_layers_per_head = getattr(config, "number_of_layers_per_head", 1)
973
+ freeze_backbone = (
974
+ freeze_backbone
975
+ if freeze_backbone is not None
976
+ else getattr(config, "freeze_backbone", False)
977
+ )
978
+ self.num_frozen_encoders = getattr(config, "num_frozen_encoders", 0)
979
+ classifier_dropout = getattr(config, "classifier_dropout", 0.1)
980
+
981
+ # Class weights for loss computation
982
+ self.use_class_weights = getattr(config, "use_class_weights", False)
983
+ self.class_weights = getattr(config, "class_weights", None)
984
+
985
+ # Initialize the transformer backbone
986
+ if base_model is None:
987
+ from transformers import AutoConfig, RobertaModel
988
+
989
+ # Use backbone_model_name if available, fallback to name_or_path
990
+ backbone_name = getattr(config, "backbone_model_name", None)
991
+ if backbone_name is None:
992
+ backbone_name = getattr(config, "name_or_path", None) or getattr(
993
+ config, "_name_or_path", None
994
+ )
995
+
996
+ if backbone_name:
997
+ # Load pretrained weights
998
+ backbone_config = AutoConfig.from_pretrained(backbone_name)
999
+ backbone_config.hidden_dropout_prob = getattr(
1000
+ config, "hidden_dropout_prob", 0.1
1001
+ )
1002
+ self.roberta = RobertaModel.from_pretrained(
1003
+ backbone_name, config=backbone_config, add_pooling_layer=False
1004
+ )
1005
+ # Store backbone_model_name in config for future loading
1006
+ if (
1007
+ not hasattr(config, "backbone_model_name")
1008
+ or config.backbone_model_name is None
1009
+ ):
1010
+ config.backbone_model_name = backbone_name
1011
+ else:
1012
+ # Fallback: initialize without pretrained weights (not recommended)
1013
+ self.roberta = RobertaModel(config, add_pooling_layer=False)
1014
+ else:
1015
+ if hasattr(base_model, "roberta"):
1016
+ self.roberta = base_model.roberta
1017
+ else:
1018
+ self.roberta = base_model
1019
+
1020
+ self.hidden_size = config.hidden_size
1021
+ self.dropout = nn.Dropout(getattr(config, "hidden_dropout_prob", 0.1))
1022
+
1023
+ # Create heads for each entity type
1024
+ print(
1025
+ f"Creating Multi-Head model with entity types: {sorted(self.entity_types)}"
1026
+ )
1027
+
1028
+ for entity_type in self.entity_types:
1029
+ # Dense layers per head
1030
+ for i in range(self.number_of_layers_per_head):
1031
+ setattr(
1032
+ self,
1033
+ f"{entity_type}_dense_{i}",
1034
+ nn.Linear(self.hidden_size, self.hidden_size),
1035
+ )
1036
+ setattr(
1037
+ self,
1038
+ f"{entity_type}_dense_activation_{i}",
1039
+ nn.GELU(approximate="none"),
1040
+ )
1041
+ setattr(
1042
+ self, f"{entity_type}_dropout_{i}", nn.Dropout(classifier_dropout)
1043
+ )
1044
+
1045
+ # Classifier per head (no CRF)
1046
+ setattr(
1047
+ self,
1048
+ f"{entity_type}_classifier",
1049
+ nn.Linear(self.hidden_size, self.num_labels),
1050
+ )
1051
+
1052
+ # Set up loss functions per entity type (with optional class weights)
1053
+ self.loss_fns = nn.ModuleDict()
1054
+ for entity_type in self.entity_types:
1055
+ if (
1056
+ self.use_class_weights
1057
+ and self.class_weights
1058
+ and entity_type in self.class_weights
1059
+ ):
1060
+ weight = torch.tensor(
1061
+ self.class_weights[entity_type], dtype=torch.float
1062
+ )
1063
+ self.loss_fns[entity_type] = nn.CrossEntropyLoss(
1064
+ weight=weight, ignore_index=-100
1065
+ )
1066
+ else:
1067
+ self.loss_fns[entity_type] = nn.CrossEntropyLoss(ignore_index=-100)
1068
+
1069
+ # Handle freezing
1070
+ if freeze_backbone:
1071
+ self._freeze_backbone()
1072
+
1073
+ def _freeze_backbone(self):
1074
+ """Freeze transformer backbone parameters."""
1075
+ print("+" * 30, "\n\n", "Freezing backbone...", "+" * 30, "\n\n")
1076
+
1077
+ # Freeze embeddings
1078
+ for param in self.roberta.embeddings.parameters():
1079
+ param.requires_grad = False
1080
+
1081
+ # Optionally freeze some encoder layers
1082
+ if self.num_frozen_encoders > 0:
1083
+ for _, param in islice(
1084
+ self.roberta.encoder.named_parameters(),
1085
+ self.num_frozen_encoders * NUM_PER_LAYER,
1086
+ ):
1087
+ param.requires_grad = False
1088
+
1089
+ def reset_head_parameters(self):
1090
+ """Reset parameters for all heads (useful after loading pretrained weights)."""
1091
+ for entity_type in self.entity_types:
1092
+ for i in range(self.number_of_layers_per_head):
1093
+ getattr(self, f"{entity_type}_dense_{i}").reset_parameters()
1094
+ getattr(self, f"{entity_type}_classifier").reset_parameters()
1095
+
1096
+ def forward(
1097
+ self,
1098
+ input_ids: Optional[torch.LongTensor] = None,
1099
+ attention_mask: Optional[torch.FloatTensor] = None,
1100
+ token_type_ids: Optional[torch.LongTensor] = None,
1101
+ position_ids: Optional[torch.LongTensor] = None,
1102
+ head_mask: Optional[torch.FloatTensor] = None,
1103
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1104
+ labels: Optional[Dict[str, torch.LongTensor]] = None,
1105
+ output_attentions: Optional[bool] = None,
1106
+ output_hidden_states: Optional[bool] = None,
1107
+ return_dict: Optional[bool] = None,
1108
+ **kwargs,
1109
+ ):
1110
+ """
1111
+ Forward pass through the multi-head model.
1112
+
1113
+ Args:
1114
+ input_ids: Input token IDs
1115
+ attention_mask: Attention mask
1116
+ labels: Dictionary mapping entity types to label tensors
1117
+ e.g., {"DRUG": tensor, "DISEASE": tensor}
1118
+ ... other standard transformer arguments
1119
+
1120
+ Returns:
1121
+ During training (labels provided):
1122
+ Tuple of (total_loss, logits_dict) where logits_dict maps entity types to logits
1123
+ During inference (no labels):
1124
+ List of prediction tensors, one per entity type (sorted alphabetically)
1125
+ """
1126
+ return_dict = (
1127
+ return_dict if return_dict is not None else self.config.use_return_dict
1128
+ )
1129
+
1130
+ # Get transformer outputs
1131
+ outputs = self.roberta(
1132
+ input_ids,
1133
+ attention_mask=attention_mask,
1134
+ token_type_ids=token_type_ids,
1135
+ position_ids=position_ids,
1136
+ head_mask=head_mask,
1137
+ inputs_embeds=inputs_embeds,
1138
+ output_attentions=output_attentions,
1139
+ output_hidden_states=output_hidden_states,
1140
+ return_dict=return_dict,
1141
+ )
1142
+
1143
+ sequence_output = outputs[0]
1144
+ sequence_output = self.dropout(sequence_output) # (batch, seq_len, hidden)
1145
+
1146
+ # Compute logits for each head
1147
+ logits = {}
1148
+ for entity_type in self.entity_types:
1149
+ head_output = sequence_output
1150
+ for i in range(self.number_of_layers_per_head):
1151
+ head_output = getattr(self, f"{entity_type}_dense_{i}")(head_output)
1152
+ head_output = getattr(self, f"{entity_type}_dense_activation_{i}")(
1153
+ head_output
1154
+ )
1155
+ head_output = getattr(self, f"{entity_type}_dropout_{i}")(head_output)
1156
+ logits[entity_type] = getattr(self, f"{entity_type}_classifier")(
1157
+ head_output
1158
+ )
1159
+
1160
+ if labels is not None:
1161
+ # Training mode - compute CrossEntropyLoss for each head
1162
+ losses = {}
1163
+
1164
+ for entity_type in self.entity_types:
1165
+ if entity_type in labels:
1166
+ entity_labels = (
1167
+ labels[entity_type].long().to(logits[entity_type].device)
1168
+ )
1169
+ entity_logits = logits[entity_type]
1170
+
1171
+ # Reshape for CrossEntropyLoss: (batch * seq_len, num_labels) and (batch * seq_len,)
1172
+ loss_fct = self.loss_fns[entity_type]
1173
+
1174
+ # Move loss function weights to the same device if needed
1175
+ if hasattr(loss_fct, "weight") and loss_fct.weight is not None:
1176
+ loss_fct.weight = loss_fct.weight.to(entity_logits.device)
1177
+
1178
+ losses[entity_type] = loss_fct(
1179
+ entity_logits.view(-1, self.num_labels),
1180
+ entity_labels.view(-1),
1181
+ )
1182
+
1183
+ # Sum losses from all heads
1184
+ total_loss = sum(losses.values())
1185
+ return total_loss, logits
1186
+
1187
+ else:
1188
+ # Inference mode - argmax decoding for each head
1189
+ predictions = {}
1190
+
1191
+ for entity_type in self.entity_types:
1192
+ # Simple argmax decoding (no CRF constraints)
1193
+ preds = torch.argmax(logits[entity_type], dim=-1)
1194
+ predictions[entity_type] = preds
1195
+
1196
+ # Return as list sorted by entity type for consistency
1197
+ return [predictions[ent] for ent in sorted(self.entity_types)]
1198
+
1199
+ def get_input_embeddings(self):
1200
+ return self.roberta.get_input_embeddings()
1201
+
1202
+ def set_input_embeddings(self, value):
1203
+ self.roberta.set_input_embeddings(value)
1204
+
1205
+ @classmethod
1206
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
1207
+ """Override from_pretrained to handle custom model loading."""
1208
+ import json
1209
+ import os
1210
+
1211
+ config = kwargs.pop("config", None)
1212
+
1213
+ if config is None:
1214
+ # Load config directly from JSON to get all saved attributes
1215
+ config_file = os.path.join(pretrained_model_name_or_path, "config.json")
1216
+ if os.path.exists(config_file):
1217
+ with open(config_file, "r") as f:
1218
+ config_dict = json.load(f)
1219
+
1220
+ # Create MultiHeadConfig with all loaded parameters
1221
+ config = MultiHeadConfig(**config_dict)
1222
+ else:
1223
+ from transformers import AutoConfig
1224
+
1225
+ config = AutoConfig.from_pretrained(
1226
+ pretrained_model_name_or_path,
1227
+ trust_remote_code=kwargs.get("trust_remote_code", True),
1228
+ )
1229
+
1230
+ # Ensure config has all required RoBERTa parameters
1231
+ roberta_defaults = {
1232
+ "layer_norm_eps": 1e-5,
1233
+ "hidden_size": 768,
1234
+ "num_hidden_layers": 12,
1235
+ "num_attention_heads": 12,
1236
+ "intermediate_size": 3072,
1237
+ "hidden_act": "gelu",
1238
+ "hidden_dropout_prob": 0.1,
1239
+ "attention_probs_dropout_prob": 0.1,
1240
+ "max_position_embeddings": 514,
1241
+ "type_vocab_size": 1,
1242
+ "initializer_range": 0.02,
1243
+ "vocab_size": 52000,
1244
+ "pad_token_id": 1,
1245
+ "bos_token_id": 0,
1246
+ "eos_token_id": 2,
1247
+ "position_embedding_type": "absolute",
1248
+ "use_cache": True,
1249
+ "is_decoder": False,
1250
+ "add_cross_attention": False,
1251
+ "chunk_size_feed_forward": 0,
1252
+ "output_hidden_states": False,
1253
+ "output_attentions": False,
1254
+ "torchscript": False,
1255
+ "tie_word_embeddings": True,
1256
+ "return_dict": True,
1257
+ "gradient_checkpointing": False,
1258
+ "pruned_heads": {},
1259
+ "problem_type": None,
1260
+ "embedding_size": None,
1261
+ }
1262
+
1263
+ for key, default_value in roberta_defaults.items():
1264
+ if not hasattr(config, key) or getattr(config, key) is None:
1265
+ setattr(config, key, default_value)
1266
+
1267
+ freeze_backbone = getattr(config, "freeze_backbone", False)
1268
+
1269
+ model = cls(config=config, freeze_backbone=freeze_backbone)
1270
+
1271
+ # Load state dict if available
1272
+ weight_file = os.path.join(pretrained_model_name_or_path, "pytorch_model.bin")
1273
+ safetensors_file = os.path.join(
1274
+ pretrained_model_name_or_path, "model.safetensors"
1275
+ )
1276
+
1277
+ try:
1278
+ if os.path.exists(safetensors_file):
1279
+ from safetensors.torch import load_file
1280
+
1281
+ state_dict = load_file(safetensors_file)
1282
+ model.load_state_dict(state_dict)
1283
+ elif os.path.exists(weight_file):
1284
+ state_dict = torch.load(weight_file, map_location="cpu")
1285
+ model.load_state_dict(state_dict)
1286
+ else:
1287
+ print(
1288
+ "Warning: No pre-trained weights found. Using randomly initialized model."
1289
+ )
1290
+ except Exception as e:
1291
+ print(f"Warning: Could not load pre-trained weights: {e}")
1292
+
1293
+ return model
1294
+
1295
+
1296
+ class TokenClassificationModel(PreTrainedModel):
1297
+ """
1298
+ Custom token classification model with configurable classifier head (no CRF).
1299
+ This model can be loaded with trust_remote_code=True for HuggingFace Hub compatibility.
1300
+ """
1301
+
1302
+ def __init__(self, config):
1303
+ super().__init__(config)
1304
+ self.config = config
1305
+ self.num_labels = config.num_labels
1306
+
1307
+ # Initialize the roberta backbone - load pretrained weights
1308
+ from transformers import AutoConfig, AutoModel
1309
+
1310
+ # Use backbone_model_name if available, fallback to name_or_path
1311
+ # This is critical because name_or_path gets overwritten during save/load
1312
+ backbone_name = getattr(config, "backbone_model_name", None)
1313
+ if backbone_name is None:
1314
+ backbone_name = getattr(config, "name_or_path", None) or getattr(
1315
+ config, "_name_or_path", None
1316
+ )
1317
+ if backbone_name is None:
1318
+ raise ValueError(
1319
+ "config.backbone_model_name (or config.name_or_path) is required to load pretrained backbone"
1320
+ )
1321
+
1322
+ # Create a clean config for the backbone
1323
+ backbone_config = AutoConfig.from_pretrained(backbone_name)
1324
+ backbone_config.hidden_dropout_prob = getattr(
1325
+ config, "hidden_dropout_prob", 0.1
1326
+ )
1327
+
1328
+ self.roberta = AutoModel.from_pretrained(
1329
+ backbone_name, config=backbone_config, add_pooling_layer=False
1330
+ )
1331
+
1332
+ # Store backbone_model_name in config for future loading
1333
+ if (
1334
+ not hasattr(config, "backbone_model_name")
1335
+ or config.backbone_model_name is None
1336
+ ):
1337
+ config.backbone_model_name = backbone_name
1338
+ self.dropout = nn.Dropout(
1339
+ config.hidden_dropout_prob
1340
+ if hasattr(config, "hidden_dropout_prob")
1341
+ else 0.1
1342
+ )
1343
+
1344
+ # Get classifier configuration
1345
+ classifier_hidden_layers = getattr(config, "classifier_hidden_layers", None)
1346
+ classifier_dropout = getattr(config, "classifier_dropout", 0.1)
1347
+
1348
+ # Build classifier head
1349
+ if classifier_hidden_layers is not None:
1350
+ # rebuild the MLP head
1351
+ in_size = config.hidden_size
1352
+ layers = []
1353
+ if classifier_hidden_layers:
1354
+ for h in classifier_hidden_layers:
1355
+ layers += [
1356
+ nn.Linear(in_size, h),
1357
+ nn.ReLU(),
1358
+ nn.Dropout(classifier_dropout),
1359
+ ]
1360
+ in_size = h
1361
+ layers.append(nn.Linear(in_size, config.num_labels))
1362
+ self.classifier = nn.Sequential(*layers)
1363
+ else:
1364
+ # Default single linear layer
1365
+ self.classifier = nn.Linear(config.hidden_size, config.num_labels)
1366
+
1367
+ # Initialize weights
1368
+ self.init_weights()
1369
+
1370
+ def forward(
1371
+ self,
1372
+ input_ids: Optional[torch.LongTensor] = None,
1373
+ attention_mask: Optional[torch.FloatTensor] = None,
1374
+ token_type_ids: Optional[torch.LongTensor] = None,
1375
+ position_ids: Optional[torch.LongTensor] = None,
1376
+ head_mask: Optional[torch.FloatTensor] = None,
1377
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1378
+ labels: Optional[torch.LongTensor] = None,
1379
+ output_attentions: Optional[bool] = None,
1380
+ output_hidden_states: Optional[bool] = None,
1381
+ return_dict: Optional[bool] = None,
1382
+ **kwargs,
1383
+ ) -> Union[Tuple[torch.Tensor], TokenClassifierOutput]:
1384
+ return_dict = (
1385
+ return_dict if return_dict is not None else self.config.use_return_dict
1386
+ )
1387
+
1388
+ # Run inputs through the RoBERTa backbone
1389
+ outputs = self.roberta(
1390
+ input_ids,
1391
+ attention_mask=attention_mask,
1392
+ token_type_ids=token_type_ids,
1393
+ position_ids=position_ids,
1394
+ head_mask=head_mask,
1395
+ inputs_embeds=inputs_embeds,
1396
+ output_attentions=output_attentions,
1397
+ output_hidden_states=output_hidden_states,
1398
+ return_dict=return_dict,
1399
+ )
1400
+
1401
+ sequence_output = outputs.last_hidden_state
1402
+ sequence_output = self.dropout(sequence_output)
1403
+ logits = self.classifier(sequence_output)
1404
+
1405
+ loss = None
1406
+ if labels is not None:
1407
+ loss_fct = nn.CrossEntropyLoss()
1408
+ if attention_mask is not None:
1409
+ # Only keep active parts of the sequence
1410
+ active_loss = attention_mask.view(-1) == 1
1411
+ active_logits = logits.view(-1, self.num_labels)[active_loss]
1412
+ active_labels = labels.view(-1)[active_loss]
1413
+ loss = loss_fct(active_logits, active_labels)
1414
+ else:
1415
+ loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))
1416
+
1417
+ if not return_dict:
1418
+ output = (logits,) + outputs[2:]
1419
+ return ((loss,) + output) if loss is not None else output
1420
+
1421
+ return TokenClassifierOutput(
1422
+ loss=loss,
1423
+ logits=logits,
1424
+ hidden_states=outputs.hidden_states,
1425
+ attentions=outputs.attentions,
1426
+ )
1427
+
1428
+ def get_input_embeddings(self):
1429
+ return self.roberta.get_input_embeddings()
1430
+
1431
+ def set_input_embeddings(self, value):
1432
+ self.roberta.set_input_embeddings(value)
1433
+
1434
+
1435
+ def load_custom_cardioner_multiclass_model(model_path: str, device: str = "auto"):
1436
+ """
1437
+ Utility function to easily load a custom CardioNER multiclass model.
1438
+
1439
+ Args:
1440
+ model_path: Path to the saved model directory
1441
+ device: Device to load model on ("auto", "cpu", "cuda", etc.)
1442
+
1443
+ Returns:
1444
+ tuple: (model, tokenizer, config)
1445
+ """
1446
+ # Validate model directory
1447
+ import os
1448
+
1449
+ import torch
1450
+ from transformers import AutoModelForTokenClassification, AutoTokenizer
1451
+
1452
+ required_files = ["config.json", "modeling.py", "pytorch_model.bin"]
1453
+ missing_files = [
1454
+ f for f in required_files if not os.path.exists(os.path.join(model_path, f))
1455
+ ]
1456
+
1457
+ if missing_files:
1458
+ raise FileNotFoundError(
1459
+ f"Missing required files in {model_path}: {missing_files}"
1460
+ )
1461
+
1462
+ print(f"Loading custom CardioNER multiclass model from: {model_path}")
1463
+
1464
+ # Load tokenizer
1465
+ tokenizer = AutoTokenizer.from_pretrained(model_path)
1466
+
1467
+ # Load model with trust_remote_code=True
1468
+ model = AutoModelForTokenClassification.from_pretrained(
1469
+ model_path, trust_remote_code=True
1470
+ )
1471
+
1472
+ # Set device
1473
+ if device == "auto":
1474
+ device = "cuda" if torch.cuda.is_available() else "cpu"
1475
+
1476
+ model = model.to(device)
1477
+
1478
+ print(f"Model loaded successfully on {device}")
1479
+ print(f"Model type: {type(model).__name__}")
1480
+ print(f"Number of labels: {model.num_labels}")
1481
+
1482
+ return model, tokenizer, model.config
1483
+
1484
+
1485
+ def load_custom_multihead_crf_model(model_path: str, device: str = "auto"):
1486
+ """
1487
+ Utility function to load a Multi-Head CRF model.
1488
+
1489
+ Args:
1490
+ model_path: Path to the saved model directory
1491
+ device: Device to load model on ("auto", "cpu", "cuda", etc.)
1492
+
1493
+ Returns:
1494
+ tuple: (model, tokenizer, config)
1495
+ """
1496
+ import os
1497
+
1498
+ from transformers import AutoTokenizer
1499
+
1500
+ # Validate model directory
1501
+ required_files = ["config.json", "modeling.py"]
1502
+ missing_files = [
1503
+ f for f in required_files if not os.path.exists(os.path.join(model_path, f))
1504
+ ]
1505
+
1506
+ if missing_files:
1507
+ raise FileNotFoundError(
1508
+ f"Missing required files in {model_path}: {missing_files}"
1509
+ )
1510
+
1511
+ print(f"Loading Multi-Head CRF model from: {model_path}")
1512
+
1513
+ # Load tokenizer
1514
+ tokenizer = AutoTokenizer.from_pretrained(model_path)
1515
+
1516
+ # Load config
1517
+ import json
1518
+
1519
+ with open(os.path.join(model_path, "config.json"), "r") as f:
1520
+ config_dict = json.load(f)
1521
+
1522
+ config = MultiHeadCRFConfig(**config_dict)
1523
+
1524
+ # Load model
1525
+ model = TokenClassificationModelMultiHeadCRF.from_pretrained(
1526
+ model_path, config=config
1527
+ )
1528
+
1529
+ # Set device
1530
+ if device == "auto":
1531
+ device = "cuda" if torch.cuda.is_available() else "cpu"
1532
+
1533
+ model = model.to(device)
1534
+
1535
+ print(f"Model loaded successfully on {device}")
1536
+ print(f"Model type: {type(model).__name__}")
1537
+ print(f"Entity types: {model.entity_types}")
1538
+ print(f"Number of labels per head: {model.num_labels}")
1539
+
1540
+ return model, tokenizer, model.config
1541
+
1542
+
1543
+ def validate_custom_multiclass_model_directory(model_path: str) -> dict:
1544
+ """
1545
+ Validate that a model directory contains all necessary files for custom multiclass model loading.
1546
+
1547
+ Args:
1548
+ model_path: Path to the model directory
1549
+
1550
+ Returns:
1551
+ dict: Validation results with status and details
1552
+ """
1553
+ import json
1554
+ import os
1555
+
1556
+ validation_results = {
1557
+ "valid": True,
1558
+ "errors": [],
1559
+ "warnings": [],
1560
+ "files_found": [],
1561
+ "model_info": {},
1562
+ }
1563
+
1564
+ # Required files
1565
+ required_files = {
1566
+ "config.json": "Model configuration",
1567
+ "modeling.py": "Custom model class definition",
1568
+ "pytorch_model.bin": "Model weights",
1569
+ }
1570
+
1571
+ # Optional files
1572
+ optional_files = {
1573
+ "tokenizer.json": "Tokenizer vocabulary",
1574
+ "tokenizer_config.json": "Tokenizer configuration",
1575
+ "training_args.json": "Training arguments",
1576
+ }
1577
+
1578
+ # Check required files
1579
+ for filename, description in required_files.items():
1580
+ filepath = os.path.join(model_path, filename)
1581
+ if os.path.exists(filepath):
1582
+ validation_results["files_found"].append(f"{filename} ({description})")
1583
+ else:
1584
+ validation_results["valid"] = False
1585
+ validation_results["errors"].append(
1586
+ f"Missing required file: {filename} - {description}"
1587
+ )
1588
+
1589
+ # Check optional files
1590
+ for filename, description in optional_files.items():
1591
+ filepath = os.path.join(model_path, filename)
1592
+ if os.path.exists(filepath):
1593
+ validation_results["files_found"].append(f"{filename} ({description})")
1594
+ else:
1595
+ validation_results["warnings"].append(
1596
+ f"Missing optional file: {filename} - {description}"
1597
+ )
1598
+
1599
+ # Parse config if available
1600
+ config_path = os.path.join(model_path, "config.json")
1601
+ if os.path.exists(config_path):
1602
+ try:
1603
+ with open(config_path, "r") as f:
1604
+ config = json.load(f)
1605
+
1606
+ validation_results["model_info"]["num_labels"] = config.get(
1607
+ "num_labels", "Unknown"
1608
+ )
1609
+ validation_results["model_info"]["model_type"] = config.get(
1610
+ "model_type", "Unknown"
1611
+ )
1612
+ validation_results["model_info"]["has_auto_map"] = "auto_map" in config
1613
+ validation_results["model_info"]["classifier_hidden_layers"] = config.get(
1614
+ "classifier_hidden_layers", None
1615
+ )
1616
+ validation_results["model_info"]["freeze_backbone"] = config.get(
1617
+ "freeze_backbone", None
1618
+ )
1619
+ validation_results["model_info"]["use_crf"] = (
1620
+ "TokenClassificationModelCRF" in str(config.get("architectures", []))
1621
+ )
1622
+
1623
+ if not config.get("auto_map"):
1624
+ validation_results["warnings"].append(
1625
+ "No auto_map found in config - may not load correctly with trust_remote_code=True"
1626
+ )
1627
+
1628
+ except json.JSONDecodeError as e:
1629
+ validation_results["valid"] = False
1630
+ validation_results["errors"].append(f"Invalid config.json: {str(e)}")
1631
+
1632
+ # Check modeling.py content
1633
+ modeling_path = os.path.join(model_path, "modeling.py")
1634
+ if os.path.exists(modeling_path):
1635
+ try:
1636
+ with open(modeling_path, "r") as f:
1637
+ content = f.read()
1638
+
1639
+ required_classes = [
1640
+ "TokenClassificationModel",
1641
+ "TokenClassificationModelCRF",
1642
+ ]
1643
+ missing_classes = [cls for cls in required_classes if cls not in content]
1644
+
1645
+ if missing_classes:
1646
+ validation_results["valid"] = False
1647
+ validation_results["errors"].append(
1648
+ f"modeling.py missing required classes: {missing_classes}"
1649
+ )
1650
+
1651
+ except Exception as e:
1652
+ validation_results["warnings"].append(
1653
+ f"Could not read modeling.py: {str(e)}"
1654
+ )
1655
+
1656
+ return validation_results
1657
+
1658
+
1659
+ # Register the MultiHeadCRF config for auto loading
1660
+ try:
1661
+ from transformers import AutoConfig
1662
+
1663
+ AutoConfig.register("multihead-crf-tagger", MultiHeadCRFConfig)
1664
+ except Exception:
1665
+ pass # Config may already be registered
1666
+
1667
+
1668
+ def patch_legacy_model(
1669
+ model_path: str, backbone_model_name: str, dry_run: bool = True
1670
+ ) -> bool:
1671
+ """
1672
+ Patch a legacy saved model by adding backbone_model_name to config.json.
1673
+
1674
+ Use this to fix models trained before backbone_model_name was added to the config.
1675
+
1676
+ Args:
1677
+ model_path: Path to the saved model directory
1678
+ backbone_model_name: The original backbone model name used during training
1679
+ (e.g., "CLTL/MedRoBERTa.nl", "GroNLP/bert-base-dutch-cased")
1680
+ dry_run: If True, only print what would be changed without modifying files
1681
+
1682
+ Returns:
1683
+ bool: True if patch was successful (or would be successful in dry_run mode)
1684
+
1685
+ Example:
1686
+ >>> # First, do a dry run to see what will change
1687
+ >>> patch_legacy_model("/path/to/saved/model", "CLTL/MedRoBERTa.nl", dry_run=True)
1688
+ >>> # Then apply the patch
1689
+ >>> patch_legacy_model("/path/to/saved/model", "CLTL/MedRoBERTa.nl", dry_run=False)
1690
+ """
1691
+ import json
1692
+ import os
1693
+ import shutil
1694
+
1695
+ config_path = os.path.join(model_path, "config.json")
1696
+
1697
+ if not os.path.exists(config_path):
1698
+ print(f"ERROR: config.json not found at {config_path}")
1699
+ return False
1700
+
1701
+ # Load existing config
1702
+ with open(config_path, "r") as f:
1703
+ config = json.load(f)
1704
+
1705
+ # Check if already patched
1706
+ if "backbone_model_name" in config:
1707
+ print(f"Model already has backbone_model_name: {config['backbone_model_name']}")
1708
+ if config["backbone_model_name"] == backbone_model_name:
1709
+ print("No changes needed.")
1710
+ return True
1711
+ else:
1712
+ print(f"WARNING: Existing backbone_model_name differs from provided value!")
1713
+ print(f" Existing: {config['backbone_model_name']}")
1714
+ print(f" Provided: {backbone_model_name}")
1715
+ if dry_run:
1716
+ print("Would update to new value (dry_run=True)")
1717
+ else:
1718
+ print("Updating to new value...")
1719
+
1720
+ # Add backbone_model_name
1721
+ config["backbone_model_name"] = backbone_model_name
1722
+
1723
+ if dry_run:
1724
+ print(f"\n[DRY RUN] Would patch {config_path}:")
1725
+ print(f' Adding: backbone_model_name = "{backbone_model_name}"')
1726
+ print("\nTo apply this patch, run with dry_run=False")
1727
+ return True
1728
+
1729
+ # Create backup
1730
+ backup_path = config_path + ".backup"
1731
+ shutil.copy2(config_path, backup_path)
1732
+ print(f"Created backup at {backup_path}")
1733
+
1734
+ # Write updated config
1735
+ with open(config_path, "w") as f:
1736
+ json.dump(config, f, indent=2)
1737
+
1738
+ print(f"Successfully patched {config_path}")
1739
+ print(f' Added: backbone_model_name = "{backbone_model_name}"')
1740
+
1741
+ return True
1742
+
1743
+
1744
+ def patch_multiple_models(
1745
+ model_paths: list, backbone_model_name: str, dry_run: bool = True
1746
+ ) -> dict:
1747
+ """
1748
+ Patch multiple legacy saved models at once.
1749
+
1750
+ Args:
1751
+ model_paths: List of paths to saved model directories
1752
+ backbone_model_name: The original backbone model name used during training
1753
+ dry_run: If True, only print what would be changed without modifying files
1754
+
1755
+ Returns:
1756
+ dict: Results for each model path
1757
+
1758
+ Example:
1759
+ >>> models = ["/path/to/model1", "/path/to/model2"]
1760
+ >>> patch_multiple_models(models, "CLTL/MedRoBERTa.nl", dry_run=False)
1761
+ """
1762
+ results = {}
1763
+ for path in model_paths:
1764
+ print(f"\n{'=' * 60}")
1765
+ print(f"Processing: {path}")
1766
+ print("=" * 60)
1767
+ results[path] = patch_legacy_model(path, backbone_model_name, dry_run)
1768
+
1769
+ # Summary
1770
+ print(f"\n{'=' * 60}")
1771
+ print("SUMMARY")
1772
+ print("=" * 60)
1773
+ success = sum(1 for v in results.values() if v)
1774
+ print(
1775
+ f"Successfully {'would patch' if dry_run else 'patched'}: {success}/{len(model_paths)}"
1776
+ )
1777
+
1778
+ return results