multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
bff799d verified
Raw
History Blame
4.45 kB
import spaces # MUST come before torch / any CUDA-touching import
import torch
import gradio as gr
from transformers import AutoTokenizer, AutoModelForMaskedLM, top_k_top_p_filtering
import torch.nn.functional as F
MODEL_ID = "doctolib-lab/doctobert-fr-base"
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForMaskedLM.from_pretrained(MODEL_ID).to("cuda")
model.eval()
MASK_TOKEN = tokenizer.mask_token # "<mask>"
@spaces.GPU(duration=15)
def predict_mask(text: str, top_k: int = 5):
"""Predict the most likely tokens to fill the <mask> in a French medical sentence.
Args:
text: A French medical sentence containing one <mask> token.
top_k: Number of top predictions to return.
"""
if MASK_TOKEN not in text:
return {"error": f"Le texte doit contenir le jeton {MASK_TOKEN}."}
inputs = tokenizer(text, return_tensors="pt")
input_ids = inputs["input_ids"]
mask_index = (input_ids[0] == tokenizer.mask_token_id).nonzero(as_tuple=True)[0]
if len(mask_index) == 0:
return {"error": "Aucun jeton <mask> trouvé dans le texte."}
mask_pos = mask_index[0].item()
with torch.no_grad():
logits = model(**inputs).logits
mask_logits = logits[0, mask_pos, :]
probs = F.softmax(mask_logits, dim=-1)
topk_probs, topk_indices = torch.topk(probs, top_k)
results = []
for i in range(top_k):
token_id = topk_indices[i].item()
token = tokenizer.convert_ids_to_tokens(token_id)
# Clean up the SentencePiece prefix
clean_token = token.replace("▁", "").strip()
prob = topk_probs[i].item()
results.append({"token": clean_token, "score": round(prob, 4)})
# Also build a filled-in sentence with the top prediction
top_token_id = topk_indices[0].unsqueeze(0)
filled_ids = input_ids[0].clone()
filled_ids[mask_pos] = top_token_id
filled_text = tokenizer.decode(filled_ids, skip_special_tokens=True)
return {
"predictions": results,
"best_prediction": filled_text,
}
CSS = """
#col-container { max-width: 900px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# 🏥 DoctoBERT-fr-base
**DoctoBERT-fr-base** est un encodeur médical français (RoBERTa, 111M de paramètres)
pré-entraîné sur des données biomédicales et cliniques françaises.
Démo de **fill-mask** : donnez une phrase médicale contenant `<mask>`
et le modèle prédit les jetons les plus probables.
🔗 [Modèle sur le Hub](https://huggingface.co/doctolib-lab/doctobert-fr-base)
"""
)
with gr.Row():
text_input = gr.Textbox(
label="Phrase médicale (avec <mask>)",
placeholder=f"Le patient souffre d'une {MASK_TOKEN} aiguë.",
scale=4,
)
run_btn = gr.Button("Prédire", variant="primary", scale=1)
with gr.Accordion("Paramètres avancés", open=False):
top_k_input = gr.Slider(
minimum=1, maximum=20, value=5, step=1,
label="Nombre de prédictions (top-k)",
)
best_pred_output = gr.Textbox(label="Meilleure prédiction", interactive=False)
predictions_output = gr.JSON(label="Top-k prédictions")
run_btn.click(
fn=predict_mask,
inputs=[text_input, top_k_input],
outputs=[predictions_output, best_pred_output],
api_name="predict",
)
gr.Examples(
examples=[
[f"Le patient souffre d'une {MASK_TOKEN} aiguë.", 5],
[f"Le médecin prescrit un traitement pour l'{MASK_TOKEN}.", 5],
[f"La dose recommandée est de {MASK_TOKEN} mg par jour.", 5],
[f"Le patient présente des symptômes de {MASK_TOKEN} chronique.", 5],
[f"L'examen clinique révèle une {MASK_TOKEN} au niveau du thorax.", 5],
],
inputs=[text_input, top_k_input],
outputs=[predictions_output, best_pred_output],
fn=predict_mask,
cache_examples=True,
cache_mode="lazy",
)
if __name__ == "__main__":
demo.launch(mcp_server=True)