Views
No views yet
1from transformers import BertPreTrainedModel, BertModel
2from transformers.modeling_outputs import TokenClassifierOutput
3from torch import nn
4from torch.nn import CrossEntropyLoss
5import torch
6
7from torchcrf import CRF
8from transformers import BertTokenizerFast
9from transformers import BertTokenizerFast, Trainer, TrainingArguments
10from transformers.trainer_utils import IntervalStrategy
11
12class BertCRF(BertPreTrainedModel):
13
14 _keys_to_ignore_on_load_unexpected = [r"pooler"]
15
16 def __init__(self, config):
17 super().__init__(config)
18 self.num_labels = config.num_labels
19
20 self.bert = BertModel(config, add_pooling_layer=False)
21 self.dropout = nn.Dropout(config.hidden_dropout_prob)
22 self.classifier = nn.Linear(config.hidden_size, config.num_labels)
23 self.crf = CRF(num_tags=config.num_labels, batch_first=True)
24 self.init_weights()
25
26 def forward(
27 self,
28 input_ids=None,
29 attention_mask=None,
30 token_type_ids=None,
31 position_ids=None,
32 head_mask=None,
33 inputs_embeds=None,
34 labels=None,
35 output_attentions=None,
36 output_hidden_states=None,
37 return_dict=None,
38 ):
39 r"""
40 labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
41 Labels for computing the token classification loss. Indices should be in ``[0, ..., config.num_labels -
42 1]``.
43 """
44 return_dict = return_dict if return_dict is not None else self.config.use_return_dict
45
46 outputs = self.bert(
47 input_ids,
48 attention_mask=attention_mask,
49 token_type_ids=token_type_ids,
50 position_ids=position_ids,
51 head_mask=head_mask,
52 inputs_embeds=inputs_embeds,
53 output_attentions=output_attentions,
54 output_hidden_states=output_hidden_states,
55 return_dict=return_dict,
56 )
57
58 sequence_output = outputs[0]
59 sequence_output = self.dropout(sequence_output)
60 logits = self.classifier(sequence_output)
61
62 loss = None
63 if labels is not None:
64 log_likelihood, tags = self.crf(logits, labels), self.crf.decode(logits)
65 loss = 0 - log_likelihood
66 else:
67 tags = self.crf.decode(logits)
68 tags = torch.Tensor(tags)
69
70 if not return_dict:
71 output = (tags,) + outputs[2:]
72 return ((loss,) + output) if loss is not None else output
73
74 return loss, tags1with io.open('./multilingual-pos-tagger-language-detection-indian-context-muril/label_encoder.pkl', 'rb') as f:
2 le = cloudpickle.load(f, encoding="latin-1")
3
4model = BertCRF.from_pretrained('./multilingual-pos-tagger-language-detection-indian-context-muril/', num_labels=210)
5tokenizer = BertTokenizerFast.from_pretrained('./data/muril-base-cased/')
6
7corpus='maru naam swagat che'
8inputs = tokenizer(corpus, max_length=512, padding=True, truncation=True, return_tensors='pt',
9 return_offsets_mapping=True)
10offset_mapping = inputs.pop("offset_mapping").cpu().numpy().tolist()
11
12outputs = model(**inputs)
13print(decode(outputs[1].numpy().tolist(), inputs['input_ids'].numpy().tolist(), offset_mapping, list(le.inverse_transform(list(range(209))))))
14
15##[{'words': ['maru', 'naam', 'swagat', 'che'], 'labels': ['gu_rom-PRP', 'gu_rom-NN', 'gu_rom-NNP', 'gu_rom-VAUX']}]| Types | Output |
|---|---|
| English | [{'words': ['my', 'name', 'is', 'swagat'], 'labels': ['en-DET', 'enNN', 'en-VB', 'en-NN']}] |
| Hindi | [{'words': ['मेरा', 'नाम', 'स्वागत', 'है'], 'labels': ['hi-PRP', 'hi-NN', 'hi-NNP', 'hi-VM']}] |
| Hindi Romanised | [{'words': ['mera', 'naam', 'swagat', 'hai'], 'labels': ['hi_romPRP', 'hi_rom-NN', 'hi_rom-NNP', 'hi_rom-VM']}] |
| Gujarati | [{'words': ['મારું', 'નામ', 'સ્વગત', 'છે'], 'labels': ['gu-PRP', 'guNN', 'gu-NNP', 'gu-VAUX']}] |
| Gujarati Romanised | [{'words': ['maru', 'naam', 'swagat', 'che'], 'labels': ['gu_romPRP', 'gu_rom-NN', 'gu_rom-NNP', 'gu_rom-VAUX']}] |