Spaces:
Sleeping
Sleeping
File size: 35,183 Bytes
862b49a 41341d0 862b49a cdc317a 6cababb cdc317a 6cababb 72ef639 cdc317a e7db87b 6cababb cdc317a 6cababb cdc317a 6cababb cdc317a 862b49a cdc317a 6cababb cdc317a 6cababb cdc317a e7db87b 7a0bb55 e7db87b c72538b e7db87b c72538b e7db87b c72538b e7db87b c72538b e7db87b c7f3e02 41341d0 c7f3e02 c72538b c7f3e02 c72538b c7f3e02 c72538b c7f3e02 c72538b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b cdc317a c57b8a9 cdc317a 6cababb cdc317a c7f3e02 cdc317a c7f3e02 cdc317a e7db87b cdc317a e7db87b cdc317a c72538b cdc317a c72538b cdc317a c7f3e02 cdc317a e7db87b cdc317a e7db87b cdc317a e7db87b cdc317a e7db87b cdc317a e7db87b c72538b e7db87b cdc317a c72538b cdc317a 6cababb cdc317a 6cababb cdc317a e7db87b cdc317a c72538b cdc317a c72538b cdc317a c72538b cdc317a 1649120 72ef639 52205a2 72ef639 52205a2 72ef639 52205a2 72ef639 c57b8a9 c72538b 6cababb c72538b 6cababb 6237464 6cababb 6237464 c57b8a9 c72538b 6237464 c72538b 6237464 cdc317a 6237464 c57b8a9 c72538b 6237464 c72538b 6237464 cdc317a 6237464 cdc317a 862b49a cdc317a c7f3e02 cdc317a 862b49a e7db87b c7f3e02 e7db87b 862b49a 7c74741 6cababb cdc317a c7f3e02 cdc317a 6cababb c7f3e02 6cababb cdc317a 6c9f7c0 c7f3e02 6c9f7c0 6cababb c7f3e02 6cababb cdc317a 6cababb cdc317a 6cababb cdc317a 6c9f7c0 cdc317a e7db87b 72ef639 e7db87b 72ef639 e7db87b cdc317a e7db87b c7f3e02 cdc317a e7db87b c7f3e02 6cababb c7f3e02 72ef639 c7f3e02 72ef639 c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 e7db87b c7f3e02 cdc317a c7f3e02 6cababb cdc317a e7db87b cdc317a c7f3e02 d1c85fc cdc317a e7db87b c7f3e02 cdc317a e7db87b c7f3e02 cdc317a e7db87b c7f3e02 5a128ba cdc317a e7db87b c7f3e02 5a128ba e7db87b 5a128ba c7f3e02 e7db87b c7f3e02 cdc317a d1c85fc cdc317a 6c9f7c0 cdc317a 6c9f7c0 cdc317a e7db87b cdc317a e7db87b c7f3e02 cdc317a c7f3e02 e7db87b c7f3e02 cdc317a 6cababb 862b49a c72538b c7f3e02 c72538b 862b49a cdc317a 1649120 c7f3e02 52205a2 862b49a 7c74741 c7f3e02 6c9f7c0 6cababb cdc317a 6cababb c7f3e02 6cababb c7f3e02 1649120 6cababb c7f3e02 1649120 862b49a c7f3e02 862b49a c7f3e02 cdc317a 862b49a 6cababb cdc317a 6cababb e7db87b c7f3e02 e7db87b c7f3e02 cdc317a e7db87b 5a128ba cdc317a 862b49a cdc317a e7db87b c7f3e02 cdc317a 862b49a c7f3e02 862b49a 72ef639 52205a2 6cababb cdc317a 6cababb cdc317a 6cababb 862b49a cdc317a 862b49a cdc317a 6c9f7c0 cdc317a 862b49a c72538b 7c74741 862b49a cdc317a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 | import json
import gradio as gr
import spaces
from backbone_utils import extract_all_features, get_cached_features
from classical_ml_utils import train_classical_model
from data_utils import dataset_overview, get_class_names, get_images_for_gallery
from predict_utils import predict_uploaded_image, test_random_sample
from train_utils import (
evaluate_saved_model,
list_saved_models,
model_meta_path,
saved_model_file_path,
train_cnn,
train_mlp,
)
# ---------------------------------------------------------------------------
# Tab 1 — Dataset
# ---------------------------------------------------------------------------
def load_dataset_callback():
try:
summary, distribution_df = dataset_overview()
class_names = ["Toutes les classes"] + get_class_names()
return summary, distribution_df, gr.update(choices=class_names, value="Toutes les classes")
except Exception as e:
return {"Erreur": str(e)}, None, gr.update()
def refresh_gallery_callback(split_name, class_name, max_images):
try:
return get_images_for_gallery(split_name, class_name, int(max_images))
except Exception as e:
return [(None, f"Erreur : {e}")]
# ---------------------------------------------------------------------------
# Tab 2 — MLP (baseline)
# ---------------------------------------------------------------------------
def mlp_gpu_duration(
num_layers, hidden_dim, dropout,
learning_rate, weight_decay, batch_size, epochs,
model_tag,
request: gr.Request,
):
# Calibré sur deux exécutions réelles :
# 2 couches, hidden_dim=256, epochs=30 -> 66.6s (2.22 s/époque)
# 5 couches, hidden_dim=1024, epochs=50 -> 140.8s (2.82 s/époque)
# Le jeu de données est minuscule (peu de pas par époque) : le temps est
# dominé par un overhead fixe, hidden_dim ne le fait varier que doucement
# (x4 sur hidden_dim -> seulement +27% par époque). On plafonne à 180s :
# nettement sous le quota journalier d'un compte gratuit (300s), pour
# qu'un seul entraînement au pire réglage ne consomme pas tout le quota
# du jour d'un·e étudiant·e.
per_epoch = 2.22 + 0.0008 * max(0, int(hidden_dim) - 256)
estimated = 15 + per_epoch * int(epochs)
return min(180, max(45, int(estimated * 1.4)))
@spaces.GPU(duration=mlp_gpu_duration)
def train_mlp_callback(
num_layers, hidden_dim, dropout,
learning_rate, weight_decay, batch_size, epochs,
model_tag,
request: gr.Request,
):
try:
session_id = request.session_hash
result = train_mlp(
session_id=session_id,
num_layers=int(num_layers),
hidden_dim=int(hidden_dim),
dropout=float(dropout),
learning_rate=float(learning_rate),
weight_decay=float(weight_decay),
batch_size=int(batch_size),
epochs=int(epochs),
model_tag=model_tag,
)
models = list_saved_models(session_id)
selected = result["model_name"] if result["model_name"] in models else None
return (
result["logs"],
result["history"],
result["summary"],
result["classification_report"],
result["confusion_matrix"],
result["confusion_matrix_path"],
result["loss_curve_path"],
gr.update(choices=models, value=selected),
)
except Exception as e:
return f"Échec de l'entraînement :\n{e}", None, None, None, None, None, None, gr.update()
# ---------------------------------------------------------------------------
# Tab 3 — SimpleCNN
# ---------------------------------------------------------------------------
def cnn_gpu_duration(
num_conv_blocks, base_filters, kernel_size, use_batchnorm,
dropout, fc_dim,
learning_rate, weight_decay, batch_size, epochs,
model_tag,
request: gr.Request,
):
# Calibré sur deux exécutions réelles :
# 3 blocs, filtres=32, noyau=3, epochs=30 -> 65.3s (2.18 s/époque)
# 5 blocs, filtres=128, noyau=5, epochs=50 -> 153.9s (3.08 s/époque)
# Le nombre de paramètres varie de ~130x entre ces deux essais mais le temps
# par époque seulement de 40% : sur ce jeu de données minuscule, le coût est
# dominé par un overhead fixe (chargement/augmentation), pas par les FLOPs
# du réseau — le nombre de paramètres surestimerait donc très largement.
# On interpole plutôt sur un score d'architecture simple. Plafond 180s :
# nettement sous le quota journalier d'un compte gratuit (300s).
score = int(num_conv_blocks) * int(base_filters) * (int(kernel_size) / 3)
baseline_score, worst_score = 96.0, 1066.7
frac = max(0.0, min(1.0, (score - baseline_score) / (worst_score - baseline_score)))
per_epoch = 2.18 + 0.9 * frac
estimated = 15 + per_epoch * int(epochs)
return min(180, max(45, int(estimated * 1.4)))
@spaces.GPU(duration=cnn_gpu_duration)
def train_cnn_callback(
num_conv_blocks, base_filters, kernel_size, use_batchnorm,
dropout, fc_dim,
learning_rate, weight_decay, batch_size, epochs,
model_tag,
request: gr.Request,
):
try:
session_id = request.session_hash
result = train_cnn(
session_id=session_id,
num_conv_blocks=int(num_conv_blocks),
base_filters=int(base_filters),
kernel_size=int(kernel_size),
use_batchnorm=bool(use_batchnorm),
dropout=float(dropout),
fc_dim=int(fc_dim),
learning_rate=float(learning_rate),
weight_decay=float(weight_decay),
batch_size=int(batch_size),
epochs=int(epochs),
model_tag=model_tag,
)
models = list_saved_models(session_id)
selected = result["model_name"] if result["model_name"] in models else None
return (
result["logs"],
result["history"],
result["summary"],
result["classification_report"],
result["confusion_matrix"],
result["confusion_matrix_path"],
result["loss_curve_path"],
gr.update(choices=models, value=selected),
)
except Exception as e:
return f"Échec de l'entraînement :\n{e}", None, None, None, None, None, None, gr.update()
# ---------------------------------------------------------------------------
# Tab 4 — Backbone + ML classique
# ---------------------------------------------------------------------------
# Mesuré : 352 images -> <10s. 30s laisse une marge x3 sans bloquer
# l'admission ZeroGPU quand il ne reste que peu de quota au visiteur.
@spaces.GPU(duration=30)
def extract_features_callback():
try:
_, class_names, counts = extract_all_features()
lines = [f"Extraction terminée — {len(class_names)} classes détectées"]
for split, n in counts.items():
lines.append(f" • {split} : {n} images → {n} vecteurs de 512 dimensions")
return "\n".join(lines)
except Exception as e:
return f"Erreur lors de l'extraction :\n{e}"
def on_clf_type_change(clf_type):
show = lambda t: gr.update(visible=(clf_type == t))
return show("SVM"), show("Régression logistique"), show("k-NN"), show("Forêt aléatoire")
def train_classical_callback(
clf_type,
svm_c,
logreg_c,
knn_k,
rf_n_estimators,
use_cv,
model_tag,
request: gr.Request,
):
try:
session_id = request.session_hash
features_cache = get_cached_features()
if features_cache is None:
return (
{"Erreur": "Veuillez d'abord extraire les caractéristiques (bouton ci-dessus)."},
None, None, None, gr.update(),
)
params = {}
if clf_type == "SVM":
params = {"C": float(svm_c)}
elif clf_type == "Régression logistique":
params = {"C": float(logreg_c)}
elif clf_type == "k-NN":
params = {"n_neighbors": int(knn_k)}
elif clf_type == "Forêt aléatoire":
params = {"n_estimators": int(rf_n_estimators)}
class_names = get_class_names()
result = train_classical_model(
clf_type, features_cache, class_names, session_id,
model_tag=model_tag, use_cv=bool(use_cv), **params
)
models = list_saved_models(session_id)
selected = result["model_name"] if result["model_name"] in models else None
return (
result["summary"],
result["classification_report"],
result["confusion_matrix"],
result["confusion_matrix_path"],
gr.update(choices=models, value=selected),
)
except Exception as e:
return {"Erreur": str(e)}, None, None, None, gr.update()
# ---------------------------------------------------------------------------
# Tab 5 — Tester et analyser
# ---------------------------------------------------------------------------
def refresh_models_callback(request: gr.Request):
models = list_saved_models(request.session_hash)
return gr.update(choices=models, value=models[0] if models else None)
def get_model_info_callback(model_name, request: gr.Request):
if not model_name:
return {"message": "Aucun modèle sélectionné."}
try:
with open(model_meta_path(model_name, request.session_hash), "r", encoding="utf-8") as f:
return json.load(f)
except FileNotFoundError:
return {"message": "Métadonnées introuvables."}
def download_model_callback(model_name, request: gr.Request):
if not model_name:
return None
try:
return saved_model_file_path(model_name, request.session_hash)
except FileNotFoundError:
return None
# Une seule passe forward sur le jeu de test (53 images) sur un modèle bien
# plus léger que le backbone ResNet18 (mesuré <10s sur 352 images) : quelques
# secondes réelles. 30s de marge, pas 120s, pour ne pas bloquer l'admission
# ZeroGPU quand il reste peu de quota au visiteur.
@spaces.GPU(duration=30)
def evaluate_callback(model_name, request: gr.Request):
try:
summary, report_df, cm_df, cm_path = evaluate_saved_model(model_name, request.session_hash)
return summary, report_df, cm_df, cm_path
except Exception as e:
return {"Erreur": str(e)}, None, None, None
# Une seule image, inférence pure. 20s de marge (chargement du modèle inclus).
@spaces.GPU(duration=20)
def predict_callback(model_name, image, request: gr.Request):
try:
return predict_uploaded_image(model_name, image, request.session_hash)
except Exception as e:
return f"Échec :\n{e}", None
# Une seule image, inférence pure. 20s de marge (chargement du modèle inclus).
@spaces.GPU(duration=20)
def random_test_callback(model_name, request: gr.Request):
try:
return test_random_sample(model_name, request.session_hash)
except Exception as e:
return None, f"Échec :\n{e}", None
# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------
with gr.Blocks(title="Classification d'images microscopiques") as demo:
gr.Markdown("# Classification d'images microscopiques de charbons de bois")
gr.Markdown(
"Ce parcours pédagogique suit une progression en quatre étapes : "
"**exploration des données**, **MLP de référence**, **CNN entraîné de zéro**, "
"puis **exploitation d'un backbone préentraîné avec des algorithmes classiques**. "
"L'objectif est de comprendre pourquoi la structure convolutive et l'apprentissage par "
"transfert sont si puissants, surtout quand les données sont rares."
)
with gr.Tabs():
# ------------------------------------------------------------------ #
# Tab 1 — Explorer le dataset
# ------------------------------------------------------------------ #
with gr.Tab("1. Explorer le jeu de données"):
gr.Markdown("## Comprendre le problème avant de modéliser")
gr.Markdown(
"Avant de choisir un modèle, il est essentiel de comprendre la structure du jeu de données. "
"Combien de classes ? Combien d'images par classe ? Les classes sont-elles équilibrées ? "
"Ces questions conditionnent directement les choix de modélisation."
)
load_dataset_btn = gr.Button("Charger les informations du dataset", variant="primary")
dataset_summary = gr.JSON(label="Résumé général")
class_distribution = gr.Dataframe(
label="Distribution des images par split et par classe", interactive=False
)
gr.Markdown(
"## Visualiser les images\n"
"Parcourez des exemples d'images pour vous familiariser avec les données. "
"Notez que les images microscopiques de charbons de bois peuvent être "
"visuellement très similaires d'une espèce à l'autre — ce qui rend la tâche difficile."
)
with gr.Row():
split_selector = gr.Dropdown(
choices=["train", "validation", "test"], value="train", label="Split"
)
class_selector = gr.Dropdown(
choices=["Toutes les classes"], value="Toutes les classes", label="Classe"
)
max_images = gr.Slider(minimum=4, maximum=48, value=24, step=4, label="Nombre d'images")
refresh_gallery_btn = gr.Button("Afficher des exemples")
image_gallery = gr.Gallery(label="Exemples d'images", columns=4, height=600)
# ------------------------------------------------------------------ #
# Tab 2 — MLP (modèle de référence)
# ------------------------------------------------------------------ #
with gr.Tab("2. MLP (modèle de référence)"):
gr.Markdown("## Un premier réseau de neurones : le perceptron multicouche")
gr.Markdown(
"Avant d'introduire un CNN, commençons par le modèle le plus simple : un **MLP** "
"(perceptron multicouche), entièrement connecté. Chaque image est aplatie en un long "
"vecteur de pixels — le réseau ne sait donc rien de la structure spatiale de l'image "
"(voisinage des pixels, formes, textures locales).\n\n"
"**Contexte du problème :** notre jeu de données contient 39 espèces, "
"avec seulement 8 images par espèce en moyenne, réparties en train / validation / test.\n\n"
"**Ce que cet exercice doit montrer :** observez les courbes de perte train vs validation. "
"À partir de quelle époque la courbe de validation cesse de s'améliorer (ou remonte) "
"pendant que la perte d'entraînement continue de baisser ? C'est le signe du surapprentissage. "
"Essayez de faire varier le nombre de couches et le nombre de neurones par couche : "
"vous constaterez que le surapprentissage apparaît quel que soit le choix — "
"un MLP n'exploite pas la structure de l'image et ne peut pas s'en affranchir. "
"C'est cette limite qui motive le passage au CNN."
)
with gr.Row():
with gr.Column():
gr.Markdown("#### Architecture du MLP")
mlp_num_layers = gr.Slider(
minimum=1, maximum=4, value=2, step=1,
label="Nombre de couches cachées",
)
mlp_hidden_dim = gr.Dropdown(
choices=[64, 128, 256, 512], value=256,
label="Neurones par couche cachée",
)
gr.Markdown("#### Hyperparamètres d'entraînement")
mlp_dropout = gr.Slider(
minimum=0.0, maximum=0.8, value=0.4, step=0.05,
label="Dropout",
)
mlp_lr = gr.Number(value=1e-3, label="Taux d'apprentissage")
mlp_wd = gr.Number(value=1e-4, label="Weight decay (régularisation L2)")
mlp_bs = gr.Dropdown(choices=[8, 16, 32, 64], value=16, label="Taille du batch")
mlp_epochs = gr.Slider(
minimum=1, maximum=50, value=30, step=1, label="Nombre d'époques"
)
mlp_tag = gr.Textbox(
label="Nom du modèle", placeholder="ex. mlp_2couches_256"
)
train_mlp_btn = gr.Button("Lancer l'entraînement", variant="primary")
with gr.Column():
mlp_logs = gr.Textbox(label="Journal d'entraînement", lines=20)
mlp_history = gr.JSON(label="Historique époque par époque")
mlp_summary = gr.JSON(label="Résumé final")
gr.Markdown("## Courbes de perte (train vs validation)")
mlp_loss_curve = gr.Image(label="Perte par époque", type="filepath")
gr.Markdown("## Résultats sur le jeu de test")
mlp_report = gr.Dataframe(label="Rapport de classification", interactive=False)
mlp_cm = gr.Dataframe(label="Matrice de confusion", interactive=False)
mlp_cm_img = gr.Image(label="Matrice de confusion — figure", type="filepath")
# ------------------------------------------------------------------ #
# Tab 3 — SimpleCNN de zéro
# ------------------------------------------------------------------ #
with gr.Tab("3. CNN entraîné de zéro"):
gr.Markdown("## Entraîner un réseau convolutif sans connaissances préalables")
gr.Markdown(
"Le MLP de l'onglet précédent surapprend quels que soient les hyperparamètres choisis : "
"il ne peut pas exploiter la structure spatiale des images. Construisons maintenant un "
"réseau de neurones convolutif (CNN), conçu pour capter des motifs locaux (contours, "
"textures) grâce aux filtres de convolution, et entraînons-le directement sur nos données "
"de charbons de bois. Ce réseau part de paramètres aléatoires : il ne sait rien des images "
"au départ.\n\n"
"**Contexte du problème :** notre jeu de données contient 39 espèces, "
"avec seulement 8 images par espèce en moyenne. "
"C'est extrêmement peu pour apprendre à distinguer 39 classes visuellement similaires.\n\n"
"Jouez avec les paramètres d'architecture et d'entraînement pour observer leur effet "
"sur les performances. Essayez notamment d'augmenter la complexité du réseau "
"et observez ce qui se passe."
)
with gr.Row():
with gr.Column():
gr.Markdown("#### Architecture du CNN")
num_conv_blocks = gr.Slider(
minimum=2, maximum=4, value=3, step=1,
label="Blocs convolutionnels",
info="Chaque bloc enchaîne Conv2d → BatchNorm → ReLU → MaxPool. Plus de blocs = réseau plus profond.",
)
base_filters = gr.Dropdown(
choices=[16, 32, 64], value=32,
label="Filtres du premier bloc",
info="Le nombre de filtres double à chaque bloc. 32 → 64 → 128...",
)
kernel_size = gr.Dropdown(
choices=[3, 5], value=3,
label="Taille du noyau de convolution",
info="3×3 capte les détails fins, 5×5 capte des structures plus larges.",
)
use_batchnorm = gr.Checkbox(
value=True, label="Normalisation par lots (BatchNorm)",
info="Stabilise l'entraînement et accélère la convergence.",
)
gr.Markdown("#### Hyperparamètres d'entraînement")
cnn_dropout = gr.Slider(
minimum=0.0, maximum=0.8, value=0.4, step=0.05,
label="Dropout",
info="Désactive aléatoirement des neurones pour limiter le surapprentissage.",
)
cnn_fc_dim = gr.Dropdown(
choices=[64, 128, 256, 512], value=256,
label="Dimension de la couche cachée",
)
cnn_lr = gr.Number(value=1e-3, label="Taux d'apprentissage")
cnn_wd = gr.Number(value=1e-4, label="Weight decay (régularisation L2)")
cnn_bs = gr.Dropdown(choices=[8, 16, 32, 64], value=16, label="Taille du batch")
cnn_epochs = gr.Slider(
minimum=1, maximum=50, value=30, step=1, label="Nombre d'époques"
)
cnn_tag = gr.Textbox(
label="Nom du modèle", placeholder="ex. cnn_3blocs_32filtres"
)
train_cnn_btn = gr.Button("Lancer l'entraînement", variant="primary")
with gr.Column():
cnn_logs = gr.Textbox(label="Journal d'entraînement", lines=20)
cnn_history = gr.JSON(label="Historique époque par époque")
cnn_summary = gr.JSON(label="Résumé final")
gr.Markdown("## Courbes de perte (train vs validation)")
cnn_loss_curve = gr.Image(label="Perte par époque", type="filepath")
gr.Markdown("## Résultats sur le jeu de test")
cnn_report = gr.Dataframe(label="Rapport de classification", interactive=False)
cnn_cm = gr.Dataframe(label="Matrice de confusion", interactive=False)
cnn_cm_img = gr.Image(label="Matrice de confusion — figure", type="filepath")
# ------------------------------------------------------------------ #
# Tab 4 — Backbone préentraîné + ML classique
# ------------------------------------------------------------------ #
with gr.Tab("4. Backbone préentraîné + ML classique"):
gr.Markdown("## Exploiter les connaissances d'un modèle préentraîné")
gr.Markdown(
"Face aux limites observées avec le MLP et le CNN de zéro (peu de données, beaucoup de "
"classes), une stratégie radicalement différente consiste à réutiliser un réseau déjà "
"entraîné sur d'autres images, et à s'appuyer sur les représentations qu'il a apprises.\n\n"
"### Qu'est-ce qu'un backbone ?\n"
"Un **backbone** est un réseau convolutif dont on retire la couche de classification finale. "
"Il agit comme un extracteur de caractéristiques : pour chaque image en entrée, "
"il produit un vecteur de nombres (ici **512 dimensions**) qui encode le contenu visuel "
"de l'image de façon compacte et abstraite.\n\n"
"### Quel backbone utilisons-nous ici ?\n"
"Nous utilisons un **ResNet18 avec ses poids ImageNet d'origine, sans aucun fine-tuning** "
"sur nos images de charbons de bois. Ce modèle a été préentraîné sur ImageNet "
"(1,2 million d'images, 1 000 classes) — il n'a jamais vu une image de charbon de bois. "
"L'objectif est d'observer si des représentations apprises sur des images naturelles "
"génériques transfèrent malgré tout à un domaine très différent (microscopie).\n\n"
"### Pourquoi des algorithmes classiques ensuite ?\n"
"Une fois les images transformées en vecteurs de 512 dimensions, "
"n'importe quel algorithme de classification classique peut être appliqué. "
"Ces algorithmes (SVM, régression logistique, k-NN, forêt aléatoire) sont rapides à entraîner, "
"interprétables, et ne nécessitent pas de GPU. "
"Pour chaque algorithme, un seul hyperparamètre est ajustable — les autres réglages "
"sont fixés pour rester comparables. Une option de validation croisée permet en plus "
"d'estimer la stabilité du score sur le train set. "
"Comparez leurs résultats avec ceux obtenus aux étapes précédentes (MLP, CNN)."
)
gr.Markdown("## Étape 1 — Extraction des caractéristiques")
gr.Markdown(
"Passez toutes les images du jeu de données dans le backbone. "
"Chaque image est convertie en un vecteur de 512 valeurs. "
"Cette opération est réalisée une seule fois et mise en cache."
)
extract_btn = gr.Button(
"Extraire les caractéristiques (backbone gelé)", variant="primary"
)
extract_status = gr.Textbox(label="Statut", lines=5, interactive=False)
gr.Markdown("## Étape 2 — Entraîner un classifieur sur les caractéristiques")
gr.Markdown(
"Choisissez un algorithme et ajustez ses paramètres. "
"L'entraînement est quasi-instantané car il opère sur des vecteurs, "
"sans jamais manipuler les images brutes ni utiliser le GPU."
)
with gr.Row():
with gr.Column():
clf_type = gr.Radio(
choices=["SVM", "Régression logistique", "k-NN", "Forêt aléatoire"],
value="SVM",
label="Algorithme de classification",
)
with gr.Column(visible=True) as svm_col:
gr.Markdown("#### Paramètre SVM (noyau RBF, gamma='scale' fixés)")
svm_c = gr.Number(
value=1.0, label="C — force de régularisation",
info="Une valeur faible regularise davantage (marges plus larges).",
)
with gr.Column(visible=False) as logreg_col:
gr.Markdown("#### Paramètre Régression logistique (max_iter=1000 fixé)")
logreg_c = gr.Number(value=1.0, label="C — force de régularisation")
with gr.Column(visible=False) as knn_col:
gr.Markdown("#### Paramètre k-NN (distance euclidienne fixée)")
knn_k = gr.Slider(
minimum=1, maximum=20, value=5, step=1,
label="k — nombre de voisins",
info="k=1 mémorise les données, k élevé généralise davantage.",
)
with gr.Column(visible=False) as rf_col:
gr.Markdown("#### Paramètre Forêt aléatoire (profondeur illimitée fixée)")
rf_n_estimators = gr.Slider(
minimum=10, maximum=500, value=100, step=10, label="Nombre d'arbres"
)
cv_checkbox = gr.Checkbox(
value=False,
label="Activer la validation croisée sur le train set",
info="Ajoute un score F1 macro moyenné sur plusieurs folds, en plus du score sur le jeu de test.",
)
ml_tag = gr.Textbox(
label="Nom du modèle", placeholder="ex. svm_C1"
)
train_classical_btn = gr.Button("Entraîner le classifieur", variant="primary")
with gr.Column():
ml_summary = gr.JSON(label="Résumé des métriques")
ml_report = gr.Dataframe(label="Rapport de classification", interactive=False)
ml_cm = gr.Dataframe(label="Matrice de confusion", interactive=False)
ml_cm_img = gr.Image(label="Matrice de confusion — figure", type="filepath")
# ------------------------------------------------------------------ #
# Tab 5 — Tester et analyser
# ------------------------------------------------------------------ #
with gr.Tab("5. Tester et analyser"):
gr.Markdown("## Comparer et évaluer les modèles")
gr.Markdown(
"Tous les modèles entraînés dans les onglets précédents apparaissent ici — "
"MLP, CNN de zéro comme classifieurs ML. "
"Évaluez-les sur le jeu de test, prédisez la classe d'une image importée, "
"et tirez vos conclusions sur l'apport du backbone préentraîné."
)
with gr.Row():
with gr.Column():
model_selector = gr.Dropdown(
choices=[],
value=None,
label="Modèle sauvegardé",
info="Liste propre à votre session — les modèles des autres étudiant·e·s ne sont pas visibles ici.",
)
refresh_btn = gr.Button("Actualiser la liste")
load_info_btn = gr.Button("Afficher les informations du modèle")
model_info = gr.JSON(label="Métadonnées du modèle")
download_btn = gr.Button("Préparer le fichier à télécharger")
model_download = gr.File(
label="Fichier du modèle (poids .pt ou pipeline .joblib) — cliquez sur "
"« Préparer le fichier à télécharger » puis sur la flèche de téléchargement ci-dessous",
)
with gr.Column():
evaluate_btn = gr.Button("Évaluer sur le jeu de test", variant="primary")
eval_summary = gr.JSON(label="Résumé des métriques")
eval_report = gr.Dataframe(label="Rapport de classification", interactive=False)
eval_cm = gr.Dataframe(label="Matrice de confusion", interactive=False)
eval_cm_img = gr.Image(label="Matrice de confusion — figure", type="filepath")
gr.Markdown("## Prédiction sur une image importée")
gr.Markdown(
"Importez une image microscopique de charbon de bois et observez "
"comment les différents modèles la classifient."
)
with gr.Row():
with gr.Column():
upload_image = gr.Image(type="pil", label="Image à classer")
predict_btn = gr.Button("Prédire la classe", variant="primary")
with gr.Column():
predict_text = gr.Textbox(label="Résultat de la prédiction", lines=7)
predict_probs = gr.Label(label="Probabilités par classe")
gr.Markdown("## Test sur un échantillon aléatoire du jeu de test")
gr.Markdown(
"Tirez une image au hasard dans le jeu de test et vérifiez si le modèle "
"sélectionné la classe correctement."
)
random_test_btn = gr.Button("Tirer un échantillon aléatoire")
with gr.Row():
random_img = gr.Image(type="pil", label="Image tirée")
random_text = gr.Textbox(label="Résultat", lines=7)
random_probs = gr.Label(label="Probabilités par classe")
# ---------------------------------------------------------------------- #
# Event wiring
# ---------------------------------------------------------------------- #
load_dataset_btn.click(
fn=load_dataset_callback,
inputs=None,
outputs=[dataset_summary, class_distribution, class_selector],
)
refresh_gallery_btn.click(
fn=refresh_gallery_callback,
inputs=[split_selector, class_selector, max_images],
outputs=image_gallery,
)
train_mlp_btn.click(
fn=train_mlp_callback,
inputs=[
mlp_num_layers, mlp_hidden_dim, mlp_dropout,
mlp_lr, mlp_wd, mlp_bs, mlp_epochs,
mlp_tag,
],
outputs=[
mlp_logs, mlp_history, mlp_summary,
mlp_report, mlp_cm, mlp_cm_img, mlp_loss_curve,
model_selector,
],
)
train_cnn_btn.click(
fn=train_cnn_callback,
inputs=[
num_conv_blocks, base_filters, kernel_size, use_batchnorm,
cnn_dropout, cnn_fc_dim,
cnn_lr, cnn_wd, cnn_bs, cnn_epochs,
cnn_tag,
],
outputs=[
cnn_logs, cnn_history, cnn_summary,
cnn_report, cnn_cm, cnn_cm_img, cnn_loss_curve,
model_selector,
],
)
extract_btn.click(fn=extract_features_callback, inputs=None, outputs=extract_status)
clf_type.change(
fn=on_clf_type_change,
inputs=clf_type,
outputs=[svm_col, logreg_col, knn_col, rf_col],
)
train_classical_btn.click(
fn=train_classical_callback,
inputs=[
clf_type,
svm_c,
logreg_c,
knn_k,
rf_n_estimators,
cv_checkbox,
ml_tag,
],
outputs=[ml_summary, ml_report, ml_cm, ml_cm_img, model_selector],
)
refresh_btn.click(fn=refresh_models_callback, inputs=None, outputs=model_selector)
load_info_btn.click(
fn=get_model_info_callback, inputs=model_selector, outputs=model_info
)
model_selector.change(
fn=download_model_callback, inputs=model_selector, outputs=model_download
)
download_btn.click(
fn=download_model_callback, inputs=model_selector, outputs=model_download
)
evaluate_btn.click(
fn=evaluate_callback,
inputs=model_selector,
outputs=[eval_summary, eval_report, eval_cm, eval_cm_img],
)
predict_btn.click(
fn=predict_callback,
inputs=[model_selector, upload_image],
outputs=[predict_text, predict_probs],
)
random_test_btn.click(
fn=random_test_callback,
inputs=model_selector,
outputs=[random_img, random_text, random_probs],
)
# Peuple la liste des modèles au chargement de la page, à partir de la
# session du navigateur qui vient de se connecter (cf. refresh_models_callback).
demo.load(fn=refresh_models_callback, inputs=None, outputs=model_selector)
if __name__ == "__main__":
demo.launch(ssr_mode=False)
|