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)