GABRIEL / app.py
DonSOZA's picture
Update app.py
023b198 verified
Raw
History Blame Contribute Delete
1.98 kB
import os
from fastapi import FastAPI
from transformers import DistilBertTokenizer, DistilBertForSequenceClassification, Trainer, TrainingArguments
from datasets import load_dataset
import gradio as gr
# Créer l'application FastAPI
app = FastAPI()
# Charger le dataset 'train2' directement depuis Hugging Face
dataset = load_dataset("train2")
# Charger le modèle DistilBERT pré-entraîné
model_name = "distilbert-base-uncased" # Modèle DistilBERT de base
tokenizer = DistilBertTokenizer.from_pretrained(model_name)
model = DistilBertForSequenceClassification.from_pretrained(model_name)
# Préparer le dataset (tokenisation)
def tokenize_function(examples):
return tokenizer(examples['text'], padding="max_length", truncation=True)
train_dataset = dataset["train"].map(tokenize_function, batched=True)
eval_dataset = dataset["test"].map(tokenize_function, batched=True)
# Définir les paramètres d'entraînement
training_args = TrainingArguments(
output_dir='./results', # Répertoire où le modèle sera sauvegardé
num_train_epochs=3, # Nombre d'époques
per_device_train_batch_size=16, # Taille du batch
per_device_eval_batch_size=64, # Taille du batch pour l'évaluation
logging_dir='./logs', # Répertoire des logs
)
# Créer un Trainer
trainer = Trainer(
model=model, # Modèle DistilBERT
args=training_args, # Paramètres d'entraînement
train_dataset=train_dataset, # Dataset d'entraînement
eval_dataset=eval_dataset # Dataset d'évaluation
)
# Entraîner le modèle
trainer.train()
# Créer un pipeline Gradio pour l'interface utilisateur
def predict(text):
inputs = tokenizer(text, return_tensors="pt")
outputs = model(**inputs)
return {"prediction": outputs.logits.argmax(dim=-1).item()}
# Déployer l'application Gradio avec partage public
gr.Interface(fn=predict, inputs="text", outputs="json").launch(share=True)