run_tf_glue.py 3.92 KB
Newer Older
1
import os
thomwolf's avatar
thomwolf committed
2
3
import tensorflow as tf
import tensorflow_datasets
4
5
6
7
8
9
10
11
from transformers import (
    BertTokenizer,
    TFBertForSequenceClassification,
    BertConfig,
    glue_convert_examples_to_features,
    BertForSequenceClassification,
    glue_processors,
)
thomwolf's avatar
thomwolf committed
12

13
14
15
# script parameters
BATCH_SIZE = 32
EVAL_BATCH_SIZE = BATCH_SIZE * 2
16
17
USE_XLA = False
USE_AMP = False
Lysandre's avatar
Lysandre committed
18
19
20
21
22
23
24
25
EPOCHS = 3

TASK = "mrpc"

if TASK == "sst-2":
    TFDS_TASK = "sst2"
elif TASK == "sts-b":
    TFDS_TASK = "stsb"
26
else:
Lysandre's avatar
Lysandre committed
27
28
29
30
    TFDS_TASK = TASK

num_labels = len(glue_processors[TASK]().get_labels())
print(num_labels)
31
32
33

tf.config.optimizer.set_jit(USE_XLA)
tf.config.optimizer.set_experimental_options({"auto_mixed_precision": USE_AMP})
34

Lysandre's avatar
Lysandre committed
35
36
# Load tokenizer and model from pretrained model/vocabulary. Specify the number of labels to classify (2+: classification, 1: regression)
config = BertConfig.from_pretrained("bert-base-cased", num_labels=num_labels)
37
38
tokenizer = BertTokenizer.from_pretrained("bert-base-cased")
model = TFBertForSequenceClassification.from_pretrained("bert-base-cased", config=config)
39
40

# Load dataset via TensorFlow Datasets
41
42
data, info = tensorflow_datasets.load(f"glue/{TFDS_TASK}", with_info=True)
train_examples = info.splits["train"].num_examples
Lysandre's avatar
Lysandre committed
43
44

# MNLI expects either validation_matched or validation_mismatched
45
valid_examples = info.splits["validation"].num_examples
thomwolf's avatar
thomwolf committed
46

thomwolf's avatar
thomwolf committed
47
# Prepare dataset for GLUE as a tf.data.Dataset instance
48
train_dataset = glue_convert_examples_to_features(data["train"], tokenizer, 128, TASK)
Lysandre's avatar
Lysandre committed
49
50

# MNLI expects either validation_matched or validation_mismatched
51
valid_dataset = glue_convert_examples_to_features(data["validation"], tokenizer, 128, TASK)
52
53
train_dataset = train_dataset.shuffle(128).batch(BATCH_SIZE).repeat(-1)
valid_dataset = valid_dataset.batch(EVAL_BATCH_SIZE)
thomwolf's avatar
thomwolf committed
54

55
# Prepare training: Compile tf.keras model with optimizer, loss and learning rate schedule
56
57
58
opt = tf.keras.optimizers.Adam(learning_rate=3e-5, epsilon=1e-08)
if USE_AMP:
    # loss scaling is currently required when using mixed precision
59
    opt = tf.keras.mixed_precision.experimental.LossScaleOptimizer(opt, "dynamic")
Lysandre's avatar
Lysandre committed
60
61
62
63
64
65
66


if num_labels == 1:
    loss = tf.keras.losses.MeanSquaredError()
else:
    loss = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True)

67
metric = tf.keras.metrics.SparseCategoricalAccuracy("accuracy")
68
model.compile(optimizer=opt, loss=loss, metrics=[metric])
thomwolf's avatar
thomwolf committed
69
70

# Train and evaluate using tf.keras.Model.fit()
71
72
train_steps = train_examples // BATCH_SIZE
valid_steps = valid_examples // EVAL_BATCH_SIZE
thomwolf's avatar
thomwolf committed
73

74
75
76
77
78
79
80
history = model.fit(
    train_dataset,
    epochs=EPOCHS,
    steps_per_epoch=train_steps,
    validation_data=valid_dataset,
    validation_steps=valid_steps,
)
81
82

# Save TF2 model
83
84
os.makedirs("./save/", exist_ok=True)
model.save_pretrained("./save/")
85

86
87
if TASK == "mrpc":
    # Load the TensorFlow model in PyTorch for inspection
88
    # This is to demo the interoperability between the two frameworks, you don't have to
89
    # do this in real life (you can run the inference on the TF model).
90
    pytorch_model = BertForSequenceClassification.from_pretrained("./save/", from_tf=True)
91
92

    # Quickly test a few predictions - MRPC is a paraphrasing task, let's see if our model learned the task
93
94
95
96
97
    sentence_0 = "This research was consistent with his findings."
    sentence_1 = "His findings were compatible with this research."
    sentence_2 = "His findings were not compatible with this research."
    inputs_1 = tokenizer.encode_plus(sentence_0, sentence_1, add_special_tokens=True, return_tensors="pt")
    inputs_2 = tokenizer.encode_plus(sentence_0, sentence_2, add_special_tokens=True, return_tensors="pt")
98
99
100
101
102
103

    del inputs_1["special_tokens_mask"]
    del inputs_2["special_tokens_mask"]

    pred_1 = pytorch_model(**inputs_1)[0].argmax().item()
    pred_2 = pytorch_model(**inputs_2)[0].argmax().item()
104
105
    print("sentence_1 is", "a paraphrase" if pred_1 else "not a paraphrase", "of sentence_0")
    print("sentence_2 is", "a paraphrase" if pred_2 else "not a paraphrase", "of sentence_0")