# Source: content/notes/libraries/huggingface.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import os
os.environ["USE_TF"] = "0"  # this independent example uses only PyTorch
import torch
from transformers import BertConfig, BertForSequenceClassification
torch.manual_seed(5)
torch.set_num_threads(1)
config = BertConfig(vocab_size=32, hidden_size=16, num_hidden_layers=1,
    num_attention_heads=2, intermediate_size=24, num_labels=2,
    hidden_dropout_prob=0., attention_probs_dropout_prob=0.)
model = BertForSequenceClassification(config)
ids = torch.tensor([[2, 8, 9, 3, 0], [2, 6, 7, 8, 3]])
mask = ids.ne(0).long()
labels = torch.tensor([0, 1])
optimizer = torch.optim.AdamW(model.parameters(), lr=.001)
before = model.classifier.weight.detach().clone()
result = model(input_ids=ids, attention_mask=mask, labels=labels)
assert result.logits.shape == (2, 2) and torch.isfinite(result.loss)
result.loss.backward()
optimizer.step()
assert not torch.equal(before, model.classifier.weight)
assistant = torch.tensor([[False, False, True, True, False],
                          [False, False, False, True, True]])
lm_labels = ids.clone().masked_fill(~assistant | ~mask.bool(), -100)
assert (lm_labels[~assistant] == -100).all()
assert (lm_labels[assistant] == ids[assistant]).all()
model.eval()
with torch.inference_mode():
    probability = model(input_ids=ids, attention_mask=mask).logits.softmax(-1)
assert torch.allclose(probability.sum(-1), torch.ones(2))
print("local classifier update, output shapes and explicit response masks passed")
