| import spaces |
| 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 <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_token = token.replace("▁", "").strip() |
| prob = topk_probs[i].item() |
| results.append({"token": clean_token, "score": round(prob, 4)}) |
|
|
| |
| 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) |