Views
No views yet
roberta-base for causal span extraction
(token classification). It identifies cause and effect text spans in sentences.1from transformers import RobertaTokenizerFast, RobertaForTokenClassification
2import torch
3
4model_name = "causal-narrative/roberta-causal-span-extractor"
5tokenizer = RobertaTokenizerFast.from_pretrained(model_name, add_prefix_space=True)
6model = RobertaForTokenClassification.from_pretrained(model_name)
7
8text = "The heavy rain caused flooding in the city."
9words = text.split()
10inputs = tokenizer(words, is_split_into_words=True, return_tensors="pt",
11 truncation=True, padding=True)
12
13with torch.no_grad():
14 outputs = model(**inputs)
15 preds = torch.argmax(outputs.logits, dim=2)[0]
16
17id2label = model.config.id2label
18word_ids = tokenizer(words, is_split_into_words=True).word_ids()
19prev = None
20for wid in word_ids:
21 if wid is not None and wid != prev:
22 print(f"{words[wid]:20s} {id2label[preds[word_ids.index(wid)].item()]}")
23 prev = wid| Label | Description |
|---|---|
| O | Non-causal token |
| B-CAUSE | Beginning of cause span |
| I-CAUSE | Inside cause span |
| B-EFFECT | Beginning of effect span |
| I-EFFECT | Inside effect span |