logasanjeev's picture
Add model card with inference provider
ca4dd76 verified
|
Raw History Blame
2.56 kB
metadata
language: en
license: apache-2.0
tags:
  - text-classification
  - multi-label
  - bert
  - go_emotions
  - emotion-classification
datasets:
  - google-research-datasets/go_emotions
metrics:
  - f1
  - precision
  - recall
widget:
  - text: I’m just chilling today.
    example_title: Neutral Example
  - text: Thank you for saving my life!
    example_title: Gratitude Example
  - text: I’m nervous about my exam tomorrow.
    example_title: Nervousness Example
inference:
  parameters:
    type: text-classification
    return_type: list
    script: inference.py

GoEmotions BERT Classifier

This is a fine-tuned BERT-base-uncased model for multi-label emotion classification on the GoEmotions dataset, predicting 28 emotions (e.g., admiration, anger, joy, neutral).

Model Details

  • Architecture: BERT-base-uncased (110M parameters)
  • Training Data: GoEmotions (58k Reddit comments, 28 emotions)
  • Loss Function: Focal Loss (gamma=2)
  • Optimizer: AdamW (lr=2e-5, weight_decay=0.01)
  • Epochs: 5
  • Hardware: Kaggle T4 x2 GPUs

Performance

  • Micro F1: 0.6025 (optimized thresholds)
  • Macro F1: 0.5266
  • Precision: 0.5425
  • Recall: 0.6775
  • Hamming Loss: 0.0372
  • Avg Positive Predictions: 1.4564

Class-wise Performance (selected):

  • Gratitude: F1 0.9120
  • Love: F1 0.8032
  • Neutral: F1 0.6827
  • Nervousness: F1 0.2564
  • Relief: F1 0.2857

Usage

The model uses optimized thresholds stored in thresholds.json for predictions. Example in Python:

from transformers import BertForSequenceClassification, BertTokenizer
import torch
import json
import requests

# Load model and tokenizer
repo_id = "logasanjeev/goemotions-bert"
model = BertForSequenceClassification.from_pretrained(repo_id)
tokenizer = BertTokenizer.from_pretrained(repo_id)

# Load thresholds
thresholds_url = f"https://huggingface.co/{repo_id}/raw/main/thresholds.json"
thresholds_data = json.loads(requests.get(thresholds_url).text)
emotion_labels = thresholds_data["emotion_labels"]
thresholds = thresholds_data["thresholds"]

# Predict
text = "I’m just chilling today."
encodings = tokenizer(text, padding='max_length', truncation=True, max_length=128, return_tensors='pt')
with torch.no_grad():
    logits = torch.sigmoid(model(**encodings).logits).numpy()[0]
predictions = [(emotion_labels[i], logit) for i, (logit, thresh) in enumerate(zip(logits, thresholds)) if logit >= thresh]
print(sorted(predictions, key=lambda x: x[1], reverse=True))
# Output: [('neutral', 0.8147)]