1import torch
2import torch.nn as nn
3from transformers import PreTrainedModel, PretrainedConfig, AutoTokenizer
4
5# Define TinyTransformer model
6class TinyTransformer(nn.Module):
7 def __init__(self, vocab_size, embed_dim, num_heads, ff_dim, num_layers):
8 super().__init__()
9 self.embedding = nn.Embedding(vocab_size, embed_dim)
10 self.pos_encoding = nn.Parameter(torch.zeros(1, 512, embed_dim))
11 encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads, dim_feedforward=ff_dim, batch_first=True)
12 self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
13 self.fc = nn.Linear(embed_dim, 1)
14 self.sigmoid = nn.Sigmoid()
15
16 def forward(self, x):
17 x = self.embedding(x) + self.pos_encoding[:, :x.size(1), :]
18 x = self.transformer(x)
19 x = x.mean(dim=1) # Global average pooling
20 x = self.fc(x)
21 return self.sigmoid(x)
22
23class TinyTransformerConfig(PretrainedConfig):
24 model_type = "tiny_transformer"
25
26 def __init__(self, vocab_size=30522, embed_dim=64, num_heads=2, ff_dim=128, num_layers=4, max_position_embeddings=512, **kwargs):
27 super().__init__(**kwargs)
28 self.vocab_size = vocab_size
29 self.embed_dim = embed_dim
30 self.num_heads = num_heads
31 self.ff_dim = ff_dim
32 self.num_layers = num_layers
33 self.max_position_embeddings = max_position_embeddings
34
35class TinyTransformerForSequenceClassification(PreTrainedModel):
36 config_class = TinyTransformerConfig
37
38 def __init__(self, config):
39 super().__init__(config)
40 self.num_labels = 1
41 self.transformer = TinyTransformer(
42 config.vocab_size,
43 config.embed_dim,
44 config.num_heads,
45 config.ff_dim,
46 config.num_layers
47 )
48
49 def forward(self, input_ids, attention_mask=None):
50 outputs = self.transformer(input_ids)
51 return {"logits": outputs}
52
53# Load the Tiny-Toxic-Detector model and tokenizer
54def load_model_and_tokenizer():
55 device = torch.device("cpu") # Due to GPU overhead inference is faster on CPU!
56
57 # Load Tiny-toxic-detector
58 config = TinyTransformerConfig.from_pretrained("AssistantsLab/Tiny-Toxic-Detector")
59 model = TinyTransformerForSequenceClassification.from_pretrained("AssistantsLab/Tiny-Toxic-Detector", config=config).to(device)
60 tokenizer = AutoTokenizer.from_pretrained("AssistantsLab/Tiny-Toxic-Detector")
61
62 return model, tokenizer, device
63
64# Prediction function
65def predict_toxicity(text, model, tokenizer, device):
66 inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=128, padding="max_length").to(device)
67 if "token_type_ids" in inputs:
68 del inputs["token_type_ids"]
69
70 with torch.no_grad():
71 outputs = model(**inputs)
72 logits = outputs["logits"].squeeze()
73 prediction = "Toxic" if logits > 0.5 else "Not Toxic"
74 return prediction
75
76def main():
77 model, tokenizer, device = load_model_and_tokenizer()
78
79 while True:
80 print("Enter text to classify (or type 'exit' to quit):")
81 text = input()
82
83 if text.lower() == 'exit':
84 print("Exiting...")
85 break
86
87 if text:
88 prediction = predict_toxicity(text, model, tokenizer, device)
89 print(f"Prediction: {prediction}")
90 else:
91 print("No text provided. Please enter some text.")
92
93if __name__ == "__main__":
94 main()