Skip to content

Commit 35db418

Browse files
committed
type annotations
1 parent 521c374 commit 35db418

7 files changed

Lines changed: 382 additions & 339 deletions

File tree

scripts/train_atomizer.py

Lines changed: 29 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import json
22

33
import numpy as np
4+
from numpy.typing import NDArray
45
from datasets import Dataset
56
from transformers import (
67
AutoTokenizer,
@@ -10,8 +11,8 @@
1011
)
1112

1213

13-
def tokenize_and_align_labels(examples):
14-
"""Tokenize each sample and align the original token labels
14+
def tokenize_and_align_labels(examples: dict[str, list]) -> dict[str, list]:
15+
"""Tokenize each sample and align the original token labels
1516
to the new subword (tokenized) structure."""
1617

1718
tokenized_outputs = tokenizer(
@@ -23,14 +24,14 @@ def tokenize_and_align_labels(examples):
2324
max_length=200 # adjust as needed
2425
)
2526

26-
labels_aligned = []
27+
labels_aligned: list[list[int]] = []
2728
for i, labels in enumerate(examples["labels"]):
2829
# The tokenizer may split single words into multiple subwords.
2930
# We create a label list the same length as input_ids,
3031
# repeating the label for all subwords of the original token.
31-
word_ids = tokenized_outputs.word_ids(batch_index=i)
32-
label_ids = []
33-
previous_word_idx = None
32+
word_ids: list[int | None] = tokenized_outputs.word_ids(batch_index=i)
33+
label_ids: list[int] = []
34+
previous_word_idx: int | None = None
3435

3536
for word_idx in word_ids:
3637
if word_idx is None:
@@ -42,30 +43,32 @@ def tokenize_and_align_labels(examples):
4243

4344
labels_aligned.append(label_ids)
4445

45-
# We dont need offset_mapping during model training, so we remove it
46+
# We don't need offset_mapping during model training, so we remove it
4647
tokenized_outputs["offset_mapping"] = [None for _ in examples["tokens"]]
47-
48+
4849
tokenized_outputs["labels"] = labels_aligned
4950
return tokenized_outputs
5051

5152

52-
def compute_metrics(eval_pred):
53-
"""Compute accuracy at the token level (simple example).
53+
def compute_metrics(eval_pred: tuple[NDArray, NDArray]) -> dict[str, float]:
54+
"""Compute accuracy at the token level (simple example).
5455
You can also compute F1, precision, recall, etc. by ignoring
5556
the -100 special tokens."""
57+
logits: NDArray
58+
labels: NDArray
5659
logits, labels = eval_pred
57-
predictions = np.argmax(logits, axis=-1)
60+
predictions: NDArray = np.argmax(logits, axis=-1)
5861

5962
# Flatten ignoring -100
60-
true_predictions = []
61-
true_labels = []
63+
true_predictions: list[int] = []
64+
true_labels: list[int] = []
6265
for pred, lab in zip(predictions, labels):
6366
for p, l in zip(pred, lab):
6467
if l != -100: # skip special tokens
6568
true_predictions.append(p)
6669
true_labels.append(l)
6770

68-
results = accuracy_metric.compute(
71+
results: dict[str, float] = accuracy_metric.compute(
6972
references=true_labels,
7073
predictions=true_predictions
7174
)
@@ -74,24 +77,24 @@ def compute_metrics(eval_pred):
7477

7578
if __name__ == '__main__':
7679
with open("sentences.jsonl", "rt") as f:
77-
sentences = [json.loads(line) for line in f]
80+
sentences: list[dict] = [json.loads(line) for line in f]
7881

79-
dataset_dict = {
82+
dataset_dict: dict[str, list] = {
8083
"tokens": [sentence["words"] for sentence in sentences],
8184
"labels": [sentence["types"] for sentence in sentences]
8285
}
8386

84-
full_dataset = Dataset.from_dict(dataset_dict)
87+
full_dataset: Dataset = Dataset.from_dict(dataset_dict)
8588

86-
max_words = max([len(sentence["words"]) for sentence in sentences])
89+
max_words: int = max([len(sentence["words"]) for sentence in sentences])
8790

8891

89-
labels = set()
92+
labels: set[str] = set()
9093
for sentence in sentences:
9194
labels |= set(sentence["types"])
9295
print(labels)
93-
label_to_id = {label: i for i, label in enumerate(labels)}
94-
id_to_label = {i: label for label, i in label_to_id.items()}
96+
label_to_id: dict[str, int] = {label: i for i, label in enumerate(labels)}
97+
id_to_label: dict[int, str] = {i: label for label, i in label_to_id.items()}
9598

9699
dataset = full_dataset.train_test_split(test_size=0.25, seed=42)
97100
train_dataset = dataset["train"]
@@ -101,14 +104,14 @@ def compute_metrics(eval_pred):
101104
print("Num test samples: ", len(test_dataset))
102105

103106

104-
model_checkpoint = "distilbert-base-multilingual-cased"
107+
model_checkpoint: str = "distilbert-base-multilingual-cased"
105108
tokenizer = AutoTokenizer.from_pretrained(model_checkpoint, use_fast=True, add_prefix_space=True)
106109

107110
# Apply to train/test datasets
108111
train_dataset = train_dataset.map(tokenize_and_align_labels, batched=True)
109112
test_dataset = test_dataset.map(tokenize_and_align_labels, batched=True)
110113

111-
# Remove columns we dont feed directly to the model
114+
# Remove columns we don't feed directly to the model
112115
# train_dataset = train_dataset.remove_columns(["tokens", "labels"])
113116
# test_dataset = test_dataset.remove_columns(["tokens", "labels"])
114117

@@ -125,7 +128,7 @@ def compute_metrics(eval_pred):
125128

126129
accuracy_metric = evaluate.load("accuracy") # type: ignore[attr-defined]
127130

128-
training_args = TrainingArguments(
131+
training_args: TrainingArguments = TrainingArguments(
129132
output_dir="./test-roberta-token-classifier",
130133
eval_strategy="epoch",
131134
save_strategy="epoch",
@@ -139,7 +142,7 @@ def compute_metrics(eval_pred):
139142
report_to="none" # Set to "tensorboard" if you want logs
140143
)
141144

142-
trainer = Trainer(
145+
trainer: Trainer = Trainer(
143146
model=model,
144147
args=training_args,
145148
train_dataset=train_dataset,
@@ -150,7 +153,7 @@ def compute_metrics(eval_pred):
150153

151154
trainer.train()
152155

153-
results = trainer.evaluate(test_dataset) # type: ignore[arg-type]
156+
results: dict[str, float] = trainer.evaluate(test_dataset) # type: ignore[arg-type]
154157
print("Test set results:", results)
155158

156159
trainer.save_model("./token-classifier")

src/hyperbase_parser_ab/alpha.py

Lines changed: 28 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,66 +1,69 @@
11
import numpy as np
2+
from numpy.typing import NDArray
3+
from scipy.sparse import spmatrix
24
from sklearn.ensemble import RandomForestClassifier
35
from sklearn.preprocessing import OneHotEncoder
6+
from spacy.tokens import Span
47

58
from hyperbase_parser_ab.atomizer import Atomizer
69

710

811
class Alpha(object):
9-
def __init__(self, cases_str=None, use_atomizer=False):
12+
def __init__(self, cases_str: str | None = None, use_atomizer: bool = False) -> None:
1013
if use_atomizer:
11-
self.atomizer = Atomizer()
14+
self.atomizer: Atomizer | None = Atomizer()
1215
elif cases_str:
1316
self.atomizer = None
1417

15-
X = []
16-
y = []
18+
X: list[tuple[str, str, str, str, str]] = []
19+
y: list[list[str]] = []
1720

1821
for line in cases_str.strip().split('\n'):
19-
sline = line.strip()
22+
sline: str = line.strip()
2023
if len(sline) > 0:
21-
row = sline.strip().split('\t')
22-
true_value = row[0]
23-
tag = row[3]
24-
dep = row[4]
25-
hpos = row[6]
26-
hdep = row[8]
27-
pos_after = row[19]
24+
row: list[str] = sline.strip().split('\t')
25+
true_value: str = row[0]
26+
tag: str = row[3]
27+
dep: str = row[4]
28+
hpos: str = row[6]
29+
hdep: str = row[8]
30+
pos_after: str = row[19]
2831

2932
y.append([true_value])
3033
X.append((tag, dep, hpos, hdep, pos_after))
3134

3235
if len(y) > 0:
33-
self.empty = False
36+
self.empty: bool = False
3437

35-
self.encX = OneHotEncoder(handle_unknown='ignore', sparse_output=False)
38+
self.encX: OneHotEncoder = OneHotEncoder(handle_unknown='ignore', sparse_output=False)
3639
self.encX.fit(np.array(X))
37-
self.ency = OneHotEncoder(handle_unknown='ignore', sparse_output=False)
40+
self.ency: OneHotEncoder = OneHotEncoder(handle_unknown='ignore', sparse_output=False)
3841
self.ency.fit(np.array(y))
3942

40-
X_ = self.encX.transform(np.array(X))
41-
y_ = self.ency.transform(np.array(y))
43+
X_: NDArray | spmatrix = self.encX.transform(np.array(X))
44+
y_: NDArray | spmatrix = self.ency.transform(np.array(y))
4245

43-
self.clf = RandomForestClassifier(random_state=777)
46+
self.clf: RandomForestClassifier = RandomForestClassifier(random_state=777)
4447
self.clf.fit(X_, y_)
4548
else:
4649
self.empty = True
4750

48-
def predict(self, sentence, features):
51+
def predict(self, sentence: Span, features: list[tuple[str, str, str, str, str]]) -> tuple[str, ...] | list[str]:
4952
if self.atomizer:
50-
preds = self.atomizer.atomize(
53+
preds: list[tuple[str, str]] = self.atomizer.atomize(
5154
sentence=str(sentence),
5255
tokens=[str(token) for token in sentence])
53-
atom_types = [pred[1] for pred in preds]
56+
atom_types: list[str] = [pred[1] for pred in preds]
5457

5558
# force known cases
5659
for i in range(len(atom_types)):
5760
if sentence[i].pos_ == 'VERB':
5861
atom_types[i] = 'P'
5962
return atom_types
6063
else:
61-
# an empty classifier allways predicts 'C'
64+
# an empty classifier always predicts 'C'
6265
if self.empty:
6366
return tuple('C' for _ in range(len(features)))
64-
_features = self.encX.transform(np.array(features))
65-
preds = self.ency.inverse_transform(self.clf.predict(_features))
66-
return tuple(pred[0] if pred else 'C' for pred in preds)
67+
_features: NDArray | spmatrix = self.encX.transform(np.array(features))
68+
preds_arr: NDArray | spmatrix = self.ency.inverse_transform(self.clf.predict(_features))
69+
return tuple(pred[0] if pred else 'C' for pred in preds_arr)

0 commit comments

Comments
 (0)