Skip to main content

[MASK]it -Efficient pre-trained encoder adaption leveraging the [MASK]

Project description

[MASK]It -Efficient pre-trained encoder adaption leveraging the [MASK]

[MASK]It library allows to adapt pre-trained encoder transformer models for text classification leveraging the pre-training fill-mask objective. It supports multi-tasking via Multi[MASK]It extension.

Install

requires python 3.10 or above

pip install maskit-learn

Disable tokenizer parallelism to allow maskit parallelized dataset preprocessing.

import os
os.environ["TOKENIZERS_PARALLELISM"] = "false"

Set device

import torch
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")

Single Task Usage

Define task

# task and model
from transformers import AutoConfig
classes = ['happy', 'sad']
verbalizer_map = {'happy':['happy', 'fun'],
                  'sad':['sad', 'cry']}
model_name = 'google-bert/bert-base-uncased'
max_length = AutoConfig.from_pretrained(model_name).max_position_embeddings

Load dataset

from torch.utils.data import DataLoader
from maskit.dataset import MaskitDataset
texts = ['I am so happy today that I cannot stay still','I am very very sad unfortunately']
labels = [0,1]
template = '{text}. This sentence is: [MASK]'
dataset = MaskitDataset(texts=texts, 
                        labels=labels, 
                        model_name=model_name, 
                        template=template, 
                        max_length=max_length)
dataloader = DataLoader(dataset=dataset, batch_size=2)

Load pre-trained model

from maskit.model import MaskitModel
model = MaskitModel(model_name=model_name,
                    verbalizer_map=verbalizer_map)
model.to(DEVICE)
print(f'Model on: {DEVICE}')

Train

Loss function and optimizer definition

from torch.nn import CrossEntropyLoss
from torch.optim import AdamW

loss_fun = CrossEntropyLoss()
optimizer = AdamW(model.named_parameters(), 1e-5)

Training loop

model.train()
epochs = 2 
for epoch in range(epochs):
    epoch_loss = 0
    for batch in dataloader:
        # Prepare batch
        batch = {key: val.to(DEVICE) for key, val in batch.items()}
        # Forward pass
        logits = model(**batch)
        labels = batch['labels']
        loss = loss_fun(logits, labels)
        # Backpropagation
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()
        # Loss tracking
        epoch_loss += loss.item()
    print(f'epoch loss: {epoch_loss}')

Inference

model.eval()
predictions = []
labels = []
for batch in dataloader:
    # Prepare batch
    batch = {key: val.to(DEVICE) for key, val in batch.items()}
    # Forward pass
    logits = model(**batch)
    predictions.extend(logits.argmax(dim=1).tolist())
    labels.extend(batch['labels'].tolist())

for idx, text in enumerate(texts):
    print(f'Current sentence: {text}')
    print(f'Ground Truth: {classes[labels[idx]]}')
    print(f'Prediction: {classes[predictions[idx]]}')
    print('-'*50)

Multi-Task Usage

Define task

# task and model
from transformers import AutoConfig
verbalizer_map = {'sentiment': {'happy':['happy', 'fun'],'sad':['sad', 'cry']},
                    'type': {'news':['news', 'journal'], 'fiction':['fiction', 'novel']}
                    }
model_name = 'google-bert/bert-base-uncased'
max_length = AutoConfig.from_pretrained(model_name).max_position_embeddings

Load dataset

from torch.utils.data import DataLoader
from maskit.dataset import MultiMaskitDataset
texts = [
    'Our colleagues from the war zone report intensification of fights',
    'As I looked in his eyes, I fell in love',
    "Today's weather is going to be sunny and warm",
    'His stomach hurt so much that he had to leave her alone'
    ]
labels = {
    'sentiment': [1,0,0,1],
    'type': [0,1,0,1]
}
template = '{text}. This sentence is: [MASK]. The text type is: [MASK]'
task_words = {
    'sentiment': 'This sentence is:',
    'type': 'The text type is:'
}
dataset = MultiMaskitDataset(texts=texts, 
                        labels=labels, 
                        model_name=model_name, 
                        template=template, 
                        task_words=task_words,
                        max_length=max_length)
dataloader = DataLoader(dataset=dataset, batch_size=2)

Load pre-trained model

from maskit.model import MultiMaskitModel
from maskit.loss import ManualWeightedLoss
model = MultiMaskitModel(model_name=model_name,
                        verbalizer_map=verbalizer_map)
model.to(DEVICE)
print(f'Model on: {DEVICE}')
# task weights
weights = [0.5, 0.5]
loss_wrapper = loss_wrapper = ManualWeightedLoss(weights=weights)
loss_wrapper.to(DEVICE)

Train

Loss function and optimizer definition

from torch.nn import CrossEntropyLoss
from torch.optim import AdamW

loss_fun = CrossEntropyLoss()
optimizer = AdamW(list(model.named_parameters())+list(loss_wrapper.parameters()), 1e-5)

Training loop

from maskit.utils import move_to_device
model.train()
print(f"Fixed task weights: {weights}")
epochs = 2
for epoch in range(epochs):
    epoch_loss = 0.0
    for step, batch in enumerate(dataloader):
        batch = {key:move_to_device(value,DEVICE) for key, value in batch.items()}
        optimizer.zero_grad()
        logits = model(**batch)
        labels = batch['labels']
        task_losses = [loss_fun(logits[task], labels[task]) for task in verbalizer_map.keys()]
        total_loss = loss_wrapper(task_losses)
        total_loss.backward()
        optimizer.step()
        epoch_loss += total_loss.item()
    print(f"Epoch {epoch + 1}: loss = {epoch_loss:.4f}")

Inference

model.eval()
all_labels = all_preds = {key:[] for key in verbalizer_map.keys()}
for batch in dataloader:
    # Prepare batch
    batch = {key:move_to_device(value,DEVICE) for key, value in batch.items()}
    # Forward pass
    logits = model(**batch)
    for task in verbalizer_map.keys():
        all_preds[task].extend(logits[task].argmax(dim=1).tolist())
        all_labels[task].extend(batch['labels'][task].tolist())

for idx, text in enumerate(texts):
    print(f'Current sentence: {text}')
    for task in verbalizer_map.keys():
        print(f'Task {task}')
        print(f'Ground Truth: {list(verbalizer_map[task].keys())[all_labels[task][idx]]}')
        print(f'Prediction: {list(verbalizer_map[task].keys())[all_preds[task][idx]]}')
    print('-'*50)

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

maskit_learn-0.1.4.tar.gz (12.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

maskit_learn-0.1.4-py3-none-any.whl (11.7 kB view details)

Uploaded Python 3

File details

Details for the file maskit_learn-0.1.4.tar.gz.

File metadata

  • Download URL: maskit_learn-0.1.4.tar.gz
  • Upload date:
  • Size: 12.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.12.9

File hashes

Hashes for maskit_learn-0.1.4.tar.gz
Algorithm Hash digest
SHA256 a0e0cea16d05942583f99d917050652a2f2b18903a01d61da1d63747aa13f7e8
MD5 caafd182f9d1501f3b25ce2ab32d40ba
BLAKE2b-256 d3133d9ca52e582730f81b70858235dcdfbf3779a0a131fe32592c12d133799a

See more details on using hashes here.

Provenance

The following attestation bundles were made for maskit_learn-0.1.4.tar.gz:

Publisher: publish.yml on eugeniaalleva/maskit

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file maskit_learn-0.1.4-py3-none-any.whl.

File metadata

  • Download URL: maskit_learn-0.1.4-py3-none-any.whl
  • Upload date:
  • Size: 11.7 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.12.9

File hashes

Hashes for maskit_learn-0.1.4-py3-none-any.whl
Algorithm Hash digest
SHA256 37db80e56014b2b4353146677adc7981eebf113c908cb5095aa544b9d90ea976
MD5 6f1f2084b6d0df9f613ffb27c51a9037
BLAKE2b-256 d73e50f8f3a9f9247adbe5e02427e06e1de0f9db5712e7000aeec56a5b1fc047

See more details on using hashes here.

Provenance

The following attestation bundles were made for maskit_learn-0.1.4-py3-none-any.whl:

Publisher: publish.yml on eugeniaalleva/maskit

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page