""" HuggingFace-compatible wrapper for XLMEduModel. Architecture: XLM-RoBERTa-large encoder → linear classifier → CRF (torchcrf) Task: BIO token classification for situation-entity segmentation (B-EDU / I-EDU / O) Loading from the Hub: from modeling_xlmedu import XLMEduConfig, XLMEduModelHF config = XLMEduConfig.from_pretrained("your-username/your-repo") model = XLMEduModelHF.from_pretrained("your-username/your-repo", config=config) Inference: from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained("FacebookAI/xlm-roberta-large") inputs = tokenizer("Hello world.", return_tensors="pt") tag_ids = model.predict(inputs["input_ids"], inputs["attention_mask"]) tags = [model.config.id2label[i] for i in tag_ids[0]] """ from __future__ import annotations from typing import Dict, List, Optional import torch from torchcrf import CRF from transformers import AutoModel, PretrainedConfig, PreTrainedModel LABELS = {"B-EDU": 0, "I-EDU": 1, "O": 2} BIO_TAGS = [tag for tag, _ in sorted(LABELS.items(), key=lambda x: x[1])] NUM_TAGS = len(BIO_TAGS) O_IDX = LABELS["O"] class XLMEduConfig(PretrainedConfig): model_type = "xlm_edu" def __init__( self, encoder_name: str = "FacebookAI/xlm-roberta-large", num_tags: int = NUM_TAGS, label2id: Optional[Dict[str, int]] = None, id2label: Optional[Dict[int, str]] = None, **kwargs, ): super().__init__(**kwargs) self.encoder_name = encoder_name self.num_tags = num_tags self.label2id = label2id or LABELS self.id2label = id2label or {v: k for k, v in LABELS.items()} class XLMEduModelHF(PreTrainedModel): config_class = XLMEduConfig def __init__(self, config: XLMEduConfig): super().__init__(config) self.encoder = AutoModel.from_pretrained(config.encoder_name, output_hidden_states=False) self.classifier = torch.nn.Linear(self.encoder.config.hidden_size, config.num_tags) self.crf = CRF(config.num_tags, batch_first=True) def forward( self, input_ids: torch.Tensor, attention_mask: torch.Tensor, labels: Optional[torch.Tensor] = None, ): """ Returns a dict with 'loss' (if labels given) and 'logits' (emission scores). logits shape: (batch, seq_len, num_tags) """ outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask) emissions = self.classifier(outputs.last_hidden_state) loss = None if labels is not None: crf_mask = attention_mask.bool() labels_crf = labels.clone() labels_crf[labels_crf == -100] = O_IDX loss = -self.crf(emissions.float(), labels_crf, mask=crf_mask, reduction="mean") return {"loss": loss, "logits": emissions} def predict( self, input_ids: torch.Tensor, attention_mask: torch.Tensor, ) -> List[List[int]]: """Viterbi-decode the best tag sequence for each sample in the batch.""" emissions = self.forward(input_ids, attention_mask)["logits"] return self.crf.decode(emissions.float(), mask=attention_mask.bool())