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 # "" @spaces.GPU(duration=15) def predict_mask(text: str, top_k: int = 5): """Predict the most likely tokens to fill the in a French medical sentence. Args: text: A French medical sentence containing one 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 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 `` 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 )", 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)