File size: 41,382 Bytes
d27cb1d
19b5e5f
d27cb1d
 
 
 
 
19b5e5f
 
045c87e
19b5e5f
fbe34db
045c87e
19b5e5f
045c87e
19b5e5f
 
 
 
045c87e
 
19b5e5f
d27cb1d
05d2106
 
d27cb1d
4335d4f
b0c1209
d27cb1d
045c87e
19b5e5f
d27cb1d
fbe34db
0857ab8
19b5e5f
045c87e
 
 
d27cb1d
19b5e5f
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
be60cf4
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be60cf4
19b5e5f
 
 
 
d27cb1d
 
 
6ced104
 
d27cb1d
 
 
 
 
 
8a861a7
 
 
 
 
 
 
 
 
 
d27cb1d
 
8a861a7
d27cb1d
 
 
 
 
 
be60cf4
d27cb1d
be60cf4
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be60cf4
d27cb1d
19b5e5f
 
d27cb1d
 
 
 
 
 
 
 
 
19b5e5f
 
d27cb1d
be60cf4
 
 
d27cb1d
 
 
be60cf4
 
 
 
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
be60cf4
d27cb1d
be60cf4
d27cb1d
 
 
 
 
 
be60cf4
d27cb1d
 
 
 
be60cf4
d27cb1d
 
be60cf4
d27cb1d
 
 
 
be60cf4
 
19b5e5f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fbe34db
045c87e
 
fbe34db
 
 
045c87e
fbe34db
045c87e
 
fbe34db
 
be60cf4
 
19b5e5f
 
 
0857ab8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
19b5e5f
0857ab8
19b5e5f
 
 
0857ab8
 
19b5e5f
 
 
 
 
 
 
0857ab8
48549a5
 
19b5e5f
 
 
 
 
 
 
 
 
 
 
be60cf4
 
 
 
d27cb1d
be60cf4
 
d27cb1d
 
be60cf4
 
19b5e5f
 
fbe34db
d27cb1d
19b5e5f
 
 
 
d27cb1d
19b5e5f
 
4335d4f
19b5e5f
 
4335d4f
 
 
 
 
 
19b5e5f
be60cf4
d27cb1d
 
 
 
fbe34db
 
19b5e5f
fbe34db
 
 
be60cf4
 
d27cb1d
 
 
 
 
be60cf4
 
 
 
d27cb1d
be60cf4
 
 
045c87e
19b5e5f
045c87e
 
 
05d2106
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
05d2106
 
d27cb1d
 
 
 
 
 
 
 
 
 
05d2106
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7cc74a6
 
 
 
 
 
 
 
 
 
 
 
05d2106
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
 
 
 
 
 
05d2106
d27cb1d
 
 
 
 
 
05d2106
 
 
d27cb1d
 
 
 
 
 
05d2106
d27cb1d
 
 
 
 
 
 
 
 
045c87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
05d2106
 
 
 
 
 
 
 
d27cb1d
 
 
 
 
 
 
 
 
045c87e
 
 
 
 
 
 
 
 
 
 
 
19b5e5f
045c87e
 
 
 
d27cb1d
 
 
 
 
 
045c87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
045c87e
d27cb1d
045c87e
 
 
 
be60cf4
045c87e
 
 
 
 
 
 
 
 
 
 
d27cb1d
 
 
be60cf4
d27cb1d
be60cf4
d27cb1d
045c87e
 
d27cb1d
045c87e
d27cb1d
 
 
 
 
 
 
 
 
045c87e
 
 
 
 
d27cb1d
045c87e
 
 
 
 
05d2106
 
 
 
 
045c87e
d27cb1d
05d2106
045c87e
 
 
 
 
 
 
d27cb1d
05d2106
 
045c87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
 
045c87e
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
045c87e
 
 
 
 
 
d27cb1d
 
 
 
05d2106
 
d27cb1d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
045c87e
 
 
 
 
 
 
 
 
d27cb1d
045c87e
d27cb1d
045c87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
045c87e
 
d27cb1d
 
 
 
045c87e
 
 
 
 
 
 
 
 
 
 
 
d27cb1d
045c87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
19b5e5f
 
045c87e
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
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
"""SoftChart β€” custom Gradio Server app for Hugging Face Spaces.

Generate a Taiko no Tatsujin chart from any audio file, using the full system:
  - SoftChartGenerator  (main model, plan-conditioned)
  - SoftChartPlanner     (auto song-level planning)
  - SoftChartBeat        (beat/downbeat for barline anchoring)
All models load from the Hub via from_pretrained. MIT licensed.
"""

import logging
import os
import re
import shutil
import tempfile
from pathlib import Path

import gradio as gr
import numpy as np
import torch
from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles

from softchart import sc2_loader
from softchart.barscript_grid import (GridFitError, barscript_grid_from_fit,
                                      describe_grid)
from softchart.generate import generate_song, generate_song_slot, load_hf
from softchart.fonts import cjk_font_path
from softchart.grid import debias_to_grid, fit_grid_fixed_bpm, fit_grid_piecewise
from softchart.hf import SoftChartPlanner
from softchart.preview_audio import synthesize_taiko_preview
from softchart.rhythm import snap_chart
from softchart.tja import append_measure_with_gogo, gogo_measure_mask, write_tja_slots
from softchart.tja_image import render_tja_image
from softchart.vocab import FPS, HOP, N_FFT, N_MELS, SR

LOGGER = logging.getLogger("softchart.space")
STATIC_DIR = Path(__file__).with_name("static")

PLAN_REPO = os.environ.get("SC_PLAN", "JacobLinCool/softchart-planner")
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

# Selectable generators: three generations x three sizes (plus the v1.5 single
# model). All are dual (slot + time), beat- and plan-conditioned, MIT licensed.
# https://huggingface.co/collections/JacobLinCool/softchart-generators
MODELS = {
    "v1.7": "JacobLinCool/softchart-v17",
    "v1.7-small": "JacobLinCool/softchart-v17-small",
    "v1.7-tiny": "JacobLinCool/softchart-v17-tiny",
    "v1.6": "JacobLinCool/softchart-v16",
    "v1.6-small": "JacobLinCool/softchart-v16-small",
    "v1.6-tiny": "JacobLinCool/softchart-v16-tiny",
    "v1.5": "JacobLinCool/softchart-v15",
}
DEFAULT_MODEL = os.environ.get("SC_MODEL", "v1.7")

# BarScript is the experimental next-generation model. It is served through
# softchart/sc2_loader.py, which resolves a release package from $SC2_PACKAGE_DIR,
# a local directory, or its own Hub repo -- a resolution order the 1.x path does
# not share. It is deliberately absent from MODELS: that dict is the *legacy*
# routing table, every entry of which must be loadable by SoftChartGenerator.
# from_pretrained, and BarScript's architecture cannot be built that way.
SC2_MODEL = "BarScript (preview)"

MODEL_LABELS = {
    "v1.7": ("v1.7 β€” richest patterns (recommended)", "v1.7 Β· 7.9M β€” best quality"),
    "v1.7-small": ("v1.7 β€” richest patterns (recommended)", "v1.7-small Β· 3.6M"),
    "v1.7-tiny": ("v1.7 β€” richest patterns (recommended)", "v1.7-tiny Β· 1.4M β€” fastest"),
    "v1.6": ("v1.6 β€” scaling ladder", "v1.6 Β· 7.9M"),
    "v1.6-small": ("v1.6 β€” scaling ladder", "v1.6-small Β· 3.6M"),
    "v1.6-tiny": ("v1.6 β€” scaling ladder", "v1.6-tiny Β· 1.4M"),
    "v1.5": ("v1.5 β€” single model", "v1.5 Β· 7.9M"),
    SC2_MODEL: ("BarScript β€” EXPERIMENTAL PREVIEW, under evaluation",
                "BarScript (preview) β€” experimental, known defects"),
}

# Density bucket (width 0.75 notes/sec) requested per course. Measured on all
# 1 155 songs of JacobLinCool/taiko-1000-parsed: median chart-span note rate is
# 1.33 / 2.01 / 3.40 / 5.27 / 6.60 notes per second for easy / normal / hard /
# oni / ura, which buckets to 1 / 2 / 4 / 7 / 8. The first four reproduce the
# values this app already shipped; ura extends the same measurement and agrees
# with the map the rest of the repo uses (scripts/infer.py, closed_loop.py,
# eval_slot.py, gen_samples.py).
COURSE_DENS = {"easy": 1, "normal": 2, "hard": 4, "oni": 7, "ura": 8}

# TJA `COURSE:` header literals, measured rather than assumed. Of the 1 155
# songs in the corpus, 229 carry an ura chart, and every one writes its header
# as `Edit` (223) or the numeric `4` (6); the string "Ura" never appears -- it
# is the dataset's normalized course NAME. `str(course).capitalize()`, which
# both TJA writers use, would emit the invalid `COURSE:Ura`.
# softchart/tja.py is vendored library code this app does not own, so the slot
# writer's header is normalized here instead. For the four courses that shipped
# before, this map reproduces the writers' output exactly, so normalization is
# a no-op on them (asserted in scripts/test_app_wiring.py).
TJA_COURSE_HEADER = {"easy": "Easy", "normal": "Normal", "hard": "Hard",
                     "oni": "Oni", "ura": "Edit"}

# Course ids the shipped planner was TRAINED on. scripts/train_planner.py builds
# its dataset with courses=("easy","normal","hard","oni") and has no other call
# site, so embedding row 4 (ura) only ever received weight decay. Probing the
# published checkpoint confirms it: mean predicted density bucket per course id
# is 0.573 / 1.126 / 2.274 / 3.789 for the trained ids -- monotone, as expected
# -- and 1.471 for id 4, i.e. it would ask for a SPARSER plan for ura than for
# oni, when ura is the densest course in the corpus. Row 4 is therefore not used.
PLANNER_COURSE_ID = {"easy": 0, "normal": 1, "hard": 2, "oni": 3}
# Courses the planner cannot answer for, and the trained course substituted for
# them. Oni is ura's direct sibling and the highest trained density contour.
# The substitution is reported to the user, never applied silently.
PLANNER_COURSE_FALLBACK = {"ura": "oni"}

CHAR = {"don": "1", "ka": "2", "don_big": "3", "ka_big": "4",
        "roll": "5", "roll_big": "6", "balloon": "7"}
SUB = 96

_MODELS = {}  # repo id -> {"gen","slot","beat"}, loaded lazily and cached
_PLANNER = {}  # the planner is model-agnostic and shared across generators
_SC2 = {}  # the BarScript release package, loaded lazily and cached


def is_sc2_model(model_choice):
    """Whether this choice routes to BarScript instead of the 1.x Hub table."""
    return model_choice == SC2_MODEL


def available_models():
    """Selectable model names. BarScript appears only when its package resolves.

    BarScript is listed FIRST, not last.  Appended at the end it fell below the
    fold of the dropdown -- present in /api/capabilities but invisible to
    someone looking for it, which reads exactly like "it was never added".
    Its group header says EXPERIMENTAL PREVIEW and selecting it renders the
    model card's defect list, so leading with it is not a recommendation; the
    default stays v1.7.
    """
    names = []
    if sc2_loader.is_available():
        names.append(SC2_MODEL)
    names.extend(MODELS)
    return names


def get_models(model_choice=DEFAULT_MODEL, *, include_planner=False):
    repo = MODELS.get(model_choice)
    if repo is None:
        raise ValueError(
            f"unknown model {model_choice!r}; choose one of {list(MODELS)}"
        )
    if repo not in _MODELS:
        u = load_hf(repo, device=DEVICE)
        if (not getattr(u, "_dual", False) or u.beat is None
                or not getattr(u, "_has_plan", False)):
            raise RuntimeError(
                f"{repo} must provide dual generation, beat, and plan conditioning"
            )
        _MODELS[repo] = {"gen": u, "slot": u, "beat": u}
    models = dict(_MODELS[repo])
    if include_planner:
        if "plan" not in _PLANNER:
            _PLANNER["plan"] = SoftChartPlanner.from_pretrained(PLAN_REPO).to(DEVICE).eval()
        models["plan"] = _PLANNER["plan"]
    return models


def get_sc2():
    """Load (and cache) the BarScript release package.

    The log-mel front end is checked against the package's own contract on
    every load: the Space computes mel itself, and a silent drift there would
    hand the encoder a distribution it never saw.
    """
    if "sc2" not in _SC2:
        sc = sc2_loader.load(device=DEVICE)
        sc2_loader.check_preprocessing(
            sc.package_dir, sr=SR, n_fft=N_FFT, hop_length=HOP,
            n_mels=N_MELS, fmin=20.0, fmax=SR / 2)
        _SC2["sc2"] = sc
    return _SC2["sc2"]


def sc2_beat_model():
    """Beat/downbeat model for the BarScript route.

    BarScript has no beat head of its own; it needs an external BPM and downbeat to
    build its bar lattice, so the default generator's beat model supplies them.
    """
    return get_models(DEFAULT_MODEL)["beat"]


def planner_course(course):
    """``(course_the_planner_is_asked_for, was_substituted)``.

    Raises for a course with neither a trained id nor a documented substitute,
    so an unsupported course fails closed instead of landing on an arbitrary
    embedding row.
    """
    if course in PLANNER_COURSE_ID:
        return course, False
    fallback = PLANNER_COURSE_FALLBACK.get(course)
    if fallback is None:
        raise ValueError(
            f"The song planner was not trained for {course}. Turn the AI "
            "planner off, or pick another difficulty."
        )
    return fallback, True


def normalize_course_header(text, course):
    """Rewrite the TJA ``COURSE:`` header to the literal players' tools expect."""
    header = TJA_COURSE_HEADER.get(str(course).lower())
    if header is None:
        raise ValueError(f"Unsupported difficulty: {course}")
    out, n = re.subn(r"(?m)^COURSE:.*$", f"COURSE:{header}", text, count=1)
    if n != 1:
        raise RuntimeError("the generated TJA has no COURSE header to normalize")
    return out


def load_logmel(path):
    import librosa

    wav, _ = librosa.load(path, sr=SR, mono=True)
    fb = librosa.filters.mel(sr=SR, n_fft=N_FFT, n_mels=N_MELS, fmin=20.0, fmax=SR / 2)
    spec = torch.stft(torch.from_numpy(wav), N_FFT, hop_length=HOP,
                      window=torch.hann_window(N_FFT), center=True, return_complex=True)
    mel = np.log(fb @ spec.abs().pow(2).numpy() + 1e-5).astype(np.float32)
    return mel, wav


def auto_plan(mel, bpm, downbeats=None):
    T = mel.shape[1]
    dur = T / FPS
    flux = np.concatenate([[0], np.maximum(0, np.diff(mel, axis=1)).sum(0)])
    beat = 60.0 / bpm
    edges = (list(downbeats[::4]) + [dur]) if (downbeats is not None and len(downbeats) >= 2) \
        else list(np.arange(0, dur, 4 * beat)) + [dur]
    vals = [float(flux[int(a * FPS):int(b * FPS)].mean()) if int(b * FPS) > int(a * FPS) else 0.0
            for a, b in zip(edges, edges[1:])]
    if not vals:
        return None
    vals = np.array(vals)
    lo, hi = np.percentile(vals, 15), np.percentile(vals, 92)
    peak = int(np.argmax(vals))
    plan = []
    for i, (a, b) in enumerate(zip(edges, edges[1:])):
        frac = (vals[i] - lo) / max(hi - lo, 1e-6)
        d8 = int(np.clip(round(frac * 7), 0, 7))
        fl = 1 if (vals[i] <= lo and 0 < i < len(vals) - 1) else (2 if i == peak and vals[i] >= hi else 0)
        plan.append([round(a, 3), round(b, 3), d8, fl])
    return plan


def learned_plan(planner, mel, course, bpm, downbeats=None):
    dur = mel.shape[1] / FPS
    beat = 60.0 / bpm
    edges = (list(downbeats[::4]) + [dur]) if (downbeats is not None and len(downbeats) >= 2) \
        else list(np.arange(0, dur, 4 * beat)) + [dur]
    feats, spans = [], []
    for a, b in zip(edges, edges[1:]):
        seg = mel[:, int(a * FPS):int(b * FPS)]
        if seg.shape[1] < 2:
            continue
        fx = np.maximum(0, np.diff(seg, axis=1)).sum(0)
        feats.append(np.concatenate([seg.mean(1), seg.std(1), [fx.mean(), fx.std(), fx.max()]]))
        spans.append((round(float(a), 3), round(float(b), 3)))
    if not feats:
        return None
    cid = PLANNER_COURSE_ID[planner_course(course)[0]]
    x = torch.tensor(np.array(feats), dtype=torch.float32)[None].to(DEVICE)
    with torch.no_grad():
        pd, pf = planner(x, torch.tensor([cid], device=DEVICE))
    d8 = pd[0].argmax(-1).cpu().numpy()
    fl = pf[0].argmax(-1).cpu().numpy()
    return [[a, b, int(d), int(f)] for (a, b), d, f in zip(spans, d8, fl)]


def group_quantize(times, phase, grid, min_run=3):
    n = len(times)
    slots = [0] * n
    i = 0
    while i < n:
        j = i
        while j + 1 < n:
            ioi = times[j + 1] - times[j]
            ref = (times[j] - times[i]) / (j - i) if j > i else ioi
            if 0.02 < ioi < 1.2 and abs(ioi - ref) < 0.22 * max(ref, 1e-6):
                j += 1
            else:
                break
        if j - i + 1 >= min_run:
            k = max(1, int(round((times[j] - times[i]) / (j - i) / grid)))
            anchor = int(round((times[i] - phase) / grid))
            for m in range(j - i + 1):
                slots[i + m] = anchor + m * k
        else:
            for m in range(i, j + 1):
                slots[m] = int(round((times[m] - phase) / grid))
        i = j + 1
    return slots


def upload_wave_name(audio_path):
    name = os.path.basename(str(audio_path).replace("\\", "/")).strip()
    name = re.sub(r"[\x00-\x1f\x7f]+", " ", name).strip()
    return name or "song.ogg"


def output_tja_path(wave_name, course, directory):
    stem = os.path.splitext(os.path.basename(wave_name))[0].strip() or "softchart"
    stem = re.sub(r"[<>:\"/\\|?*\x00-\x1f]+", "_", stem).strip(" ._") or "softchart"
    return os.path.join(directory, f"{stem}_{course}.tja")


def write_tja(gen, bpm, title, course, level, wave, downbeats=None, grid_fit=None,
              plan=None):
    hits = sorted((h["t"], CHAR[h["type"]]) for h in gen["hits"])
    beat = 60.0 / bpm
    grid = beat / (SUB / 4)
    bias = 0.0
    if grid_fit is not None and hits:
        # authoritative fitted grid: barlines ARE the fitted downbeats.
        # De-bias the generator's systematic latency (global shift only),
        # then anchor slot 0 on the last fitted barline at/before the first note.
        times, bias = debias_to_grid([t for t, _ in hits], grid_fit["phase"], grid)
        phase = grid_fit["phase"] + float(np.floor((times[0] - grid_fit["phase"]) / (4 * beat))) * 4 * beat
        q_times = list(times)
    else:
        times = np.array([t for t, _ in hits]) if hits else np.array([0.0])
        cands = np.arange(0, beat, grid / 4)
        phase = float(cands[int(np.argmin([np.mean(np.abs(((times - o) / grid) - np.round((times - o) / grid))) for o in cands]))])
        q_times = [t for t, _ in hits]
    slot_idx = group_quantize(q_times, phase, grid)
    if slot_idx and min(slot_idx) < 0:
        # note quantized just before the anchor barline: pull back whole bars
        # so nothing is dropped (barline alignment is preserved mod SUB)
        nb = int(np.ceil(-min(slot_idx) / SUB))
        slot_idx = [s + nb * SUB for s in slot_idx]
        phase -= nb * SUB * grid
    slots = {}
    for idx, (t, ch) in zip(slot_idx, hits):
        if idx >= 0 and idx not in slots:
            slots[idx] = ch
    for sp in gen["spans"]:
        i0 = int(round((sp["t0"] - bias - phase) / grid))
        i1 = int(round((sp["t1"] - bias - phase) / grid))
        while i0 in slots:
            i0 += 1
        while i1 in slots or i1 <= i0:
            i1 += 1
        if i0 >= 0:
            slots[i0] = CHAR[sp["type"]]
            slots[i1] = "8"
    if slots and grid_fit is None:
        # Gridless anchoring: shift so the first note sits on a detected
        # downbeat if one is nearby, otherwise on the first barline.
        first_t = min(slots) * grid + phase
        anchor_t = None
        if downbeats is not None and len(downbeats):
            near = downbeats[downbeats <= first_t + 0.12]
            if len(near) and first_t - near[-1] < 4 * beat:
                anchor_t = near[-1]
        shift = int(round((anchor_t - phase) / grid)) if anchor_t is not None else min(slots)
        if shift:
            slots = {k - shift: v for k, v in slots.items()}
            phase += shift * grid
    n_meas = (max(slots) // SUB + 1) if slots else 1
    measure_starts = phase + np.arange(n_meas + 1, dtype=float) * (4 * beat)
    gogo_mask = gogo_measure_mask(plan, measure_starts, n_meas)
    lines = []
    in_gogo = False
    for m in range(n_meas):
        in_gogo = append_measure_with_gogo(
            lines,
            "".join(slots.get(m * SUB + k, "0") for k in range(SUB)) + ",",
            m, gogo_mask, in_gogo)
    if in_gogo:
        lines.append("#GOGOEND")
    balloons = [10] * sum(1 for s in gen["spans"] if s["type"] == "balloon")
    return "\n".join([
        f"TITLE:{title} (SoftChart)", f"BPM:{bpm:g}", f"WAVE:{wave}",
        f"OFFSET:{-phase:.3f}", f"COURSE:{TJA_COURSE_HEADER[course]}",
        f"LEVEL:{level}", f"BALLOON:{','.join(map(str, balloons))}" if balloons else "BALLOON:",
        "", "#START", *lines, "#END"]) + "\n"


def render_audio_plan(mel, title, course, out_path, plan=None):
    import matplotlib
    matplotlib.use("Agg")
    from matplotlib import font_manager
    import matplotlib.pyplot as plt

    font_path = cjk_font_path()
    if font_path is not None:
        font_manager.fontManager.addfont(font_path)
        matplotlib.rcParams["font.family"] = font_manager.FontProperties(fname=font_path).get_name()
    matplotlib.rcParams["axes.unicode_minus"] = False

    dur = mel.shape[1] / FPS
    fig = plt.figure(figsize=(13, 4.4 if plan else 3.2))
    gs = fig.add_gridspec(2 if plan else 1, 1,
                          height_ratios=[3.0, 1.0] if plan else [1],
                          hspace=0.14 if plan else 0.0)
    ax0 = fig.add_subplot(gs[0])
    ax0.imshow(mel, aspect="auto", origin="lower",
               cmap="magma", extent=[0, dur, 0, N_MELS])
    ax0.set_ylabel("mel")
    ax0.set_title(f"{title} β€” {course}  |  full-song mel spectrogram")
    ax0.set_xlim(0, dur)
    ax0.grid(axis="x", alpha=0.18)
    if plan:
        ax0.set_xticklabels([])
        ax1 = fig.add_subplot(gs[1], sharex=ax0)
        for a, b, d, f in plan:
            c = "#d64545" if f == 2 else ("#4a90d9" if f == 1 else "#999999")
            ax1.bar((a + b) / 2, max(d, 0.15), width=max((b - a) * 0.92, 0.01),
                    color=c, alpha=0.85)
        ax1.set_xlim(0, dur)
        ax1.set_ylim(0, 8)
        ax1.set_yticks([0, 4, 8])
        ax1.set_ylabel("plan", fontsize=8)
        ax1.set_xlabel("time (s)  β€”  plan: grey=density blue=gap red=climax")
        ax1.grid(axis="x", alpha=0.18)
    else:
        ax0.set_xlabel("time (s)")
    fig.savefig(out_path, dpi=130, bbox_inches="tight")
    plt.close(fig)
    return out_path


def sc2_bar_grid(sc2, fit, beat_fit_raw, dbs, bpm, duration_sec):
    """Bar timeline for a BarScript decode: ``(grid, anchor, offset_sec)``.

    BarScript decodes onto a bar timeline supplied by the caller, so it needs
    the audio time of the first downbeat. That cannot be guessed: a timeline
    anchored on the wrong phase puts every barline in the wrong place, and the
    chart would look like a model defect rather than a missing input. The route
    therefore fails closed when no downbeat is available.

    ``fit`` is the beat fit that PASSED its own quality gate, or None. When one
    is present its downbeats become the bar edges verbatim, so a tempo change
    the fit found survives into the decode and into the TJA's ``#BPMCHANGE``
    lines. ``fit_grid_piecewise`` returns evenly spaced downbeats for a
    constant-tempo song and for the manual-BPM refit (``fit_grid_fixed_bpm``
    trusts the user's period and estimates only the phase), so those two cases
    take the same route and simply come out uniform -- there is no separate
    uniform branch to keep in sync.

    Without a passing fit the grid is re-tiled from one scalar BPM, which is
    the behaviour this route has always had and cannot represent a tempo
    change. ``beat_fit_raw`` -- the fit that failed the gate, if there was one
    -- is handed to the builder so the fallback records WHY it happened
    (inlier fraction and RMS) instead of reporting a missing fit.
    """
    if fit is not None and len(fit.get("downbeats", [])):
        offset_sec = float(fit["downbeats"][0])
        anchor = "fitted beat grid"
    elif dbs is not None and len(dbs):
        offset_sec = float(dbs[0])
        anchor = "detected downbeats"
    else:
        raise ValueError(
            "BarScript builds its bars from a downbeat and cannot place "
            "them without one. Leave the beat grid switched on, or enter the "
            "song's BPM in Advanced settings, and try again."
        )
    try:
        grid = barscript_grid_from_fit(
            fit if fit is not None else beat_fit_raw, duration_sec,
            denom=sc2_loader.SHIP_GRID_DENOM, spec=sc2.spec,
            on_invalid="fallback", fallback_bpm=float(bpm),
            fallback_offset=offset_sec,
        )
    except GridFitError as exc:  # only reachable if the fallback itself fails
        raise RuntimeError(
            f"Could not build a bar timeline for this song: {exc}"
        ) from exc
    return grid, anchor, offset_sec


def sc2_grid_metrics(grid, grid_info, anchor, offset_sec):
    """How the bar timeline is reported to the user.

    The user's question is "was this song's tempo change handled?", so the
    answer names the route (estimated barlines vs. one uniform tempo), says
    whether the timeline that came out actually varies, and gives the fit's
    quality numbers. A fallback is never silent: it carries the builder's own
    reason string.
    """
    meta = grid_info.get("grid_meta") or {}
    quality = grid_info.get("fit_quality")
    source = grid_info.get("grid_source", "unknown")
    n_tempi = int(grid_info.get("n_distinct_bar_sec", 1))
    inlier = rms = "n/a"
    if quality is not None:
        if quality["inlier_frac"] is not None:
            inlier = f"{quality['inlier_frac']:.3f}"
        if quality["rms_ms"] is not None:
            rms = f"{quality['rms_ms']:.1f} ms"
    out = {"grid_anchor": f"{anchor}, first downbeat {offset_sec:.3f} s"}
    if source == "piecewise":
        out["bar_grid"] = describe_grid(grid)
    else:
        # describe_grid's fallback branch restates the reason at full float
        # precision, which tempo_map below already gives rounded. Only the
        # structural head is kept, so the panel says each thing once.
        out["bar_grid"] = (
            f"{grid_info['n_bars']} bars, {meta.get('meter_num', '?')}/"
            f"{meta.get('meter_den', '?')} (estimated, constant); "
            f"uniform tiling from a single BPM"
        )
    if source == "piecewise" and n_tempi > 1:
        # Deliberately NOT "the tempo change was captured". The grid follows
        # the estimated barlines, and a barline the estimator split in two is
        # indistinguishable, from here, from a bar at twice the tempo -- both
        # land as a x2 spread. Measured over 60 held-out songs, 16 were served
        # a non-uniform grid and only 8 of those actually change tempo per the
        # authored trace, 9 of the 16 spanning a factor >= 1.9. So the spread
        # is stated as the measurement it is, with the ambiguity attached.
        spread = grid_info["bar_bpm_max"] / max(grid_info["bar_bpm_min"], 1e-9)
        out["tempo_map"] = (
            f"follows the estimated barlines β€” {n_tempi} distinct bar lengths, "
            f"{grid_info['bar_bpm_min']:.1f}–{grid_info['bar_bpm_max']:.1f} BPM "
            f"across {grid_info['n_bars']} bars, spread Γ—{spread:.2f}. "
            f"A barline split in two and a genuine tempo change are not "
            f"told apart here; a spread near Γ—2 is the signature of both."
        )
    elif source == "piecewise":
        out["tempo_map"] = (
            f"constant β€” the estimated barlines are evenly spaced at "
            f"{grid_info['bar_bpm_min']:.1f} BPM, so no tempo change was found"
        )
    else:
        # The overwhelmingly common fallback is a fit that missed its own
        # quality gate, and the builder's reason string then repeats the two
        # numbers at full float precision. Those numbers are stated here
        # rounded, from ``fit_quality`` rather than by re-reading the string;
        # every OTHER reason (coverage, joint-repair budget, a malformed
        # downbeat sequence) is quoted verbatim, since nothing else carries it.
        why = (f"the beat fit did not pass its quality gate: inlier {inlier}, "
               f"RMS {rms}") if (quality is not None and not quality["ok"]) \
            else (meta.get("reason") or "no reason recorded")
        out["tempo_map"] = (
            f"flattened to one uniform {grid_info['bpm']:.1f} BPM β€” a tempo "
            f"change in this song would NOT be represented ({why})"
        )
    if quality is not None:
        out["beat_fit"] = (
            f"{'passed' if quality['ok'] else 'FAILED its quality gate'} β€” "
            f"inlier {inlier}, RMS {rms}, {quality['n_segments']} tempo "
            f"segment(s) found"
        )
    else:
        out["beat_fit"] = "no beat fit β€” bars tiled from a single BPM"
    repairs = []
    if meta.get("joint_bars_merged"):
        repairs.append(f"{meta['joint_bars_merged']} splice-joint bar(s) merged")
    if meta.get("tail_bars_extrapolated"):
        repairs.append(f"{meta['tail_bars_extrapolated']} tail bar(s) extrapolated")
    if repairs:
        out["grid_repairs"] = "; ".join(repairs)
    return out


def generate_sc2(sc2, mel, fit, beat_fit_raw, dbs, bpm, course, level,
                 wave_name, title, *, sampling, temperature, top_p, seed=0):
    """Run the BarScript release. Returns ``(generated, tja, extra_metrics)``.

    Sampling parameters are left at the package's recorded serving contract
    unless the user turns creative sampling on. Anything overridden is reported
    back, because the model card's measurements describe the recorded values.
    """
    duration_sec = mel.shape[1] / FPS
    grid, anchor, offset_sec = sc2_bar_grid(
        sc2, fit, beat_fit_raw, dbs, bpm, duration_sec)

    overrides = {}
    if sampling:
        overrides = {"greedy": False, "temperature": float(temperature),
                     "top_p": float(top_p)}
    result = sc2.generate(
        mel, course, level=level, grid=grid,
        density_bucket=COURSE_DENS[course], seed=seed,
        grid_denom=sc2_loader.SHIP_GRID_DENOM,
        title=title, wave=wave_name, **overrides,
    )
    grid_info = result["grid"]
    extra = {
        # the bar-LOCAL lattice: how finely a note may sit inside one bar.
        # Unrelated to the bar timeline reported by tempo_map below.
        "chart_grid": (f"/{grid_info['denom']} per bar"
                       + ("" if grid_info["triplets_representable"]
                          else " β€” triplets impossible")),
        "density_bucket": result["density_bucket"],
        "sampling": ("user override: " + ", ".join(
            f"{k}={v}" for k, v in sorted(result["serving_overrides"].items()))
            ) if result["serving_overrides"] else "package defaults",
        **sc2_grid_metrics(grid, grid_info, anchor, offset_sec),
    }
    if not result["motif"]["verified"]:
        raise RuntimeError(
            "BarScript decoded without the motif constraint the release "
            "package resolved; refusing to present the chart."
        )
    return result["gen"], result["tja"], extra


def _progress(stage, fraction, title, detail):
    return {
        "kind": "progress",
        "stage": stage,
        "progress": fraction,
        "title": title,
        "detail": detail,
    }


def _uploaded_file(value, original_name):
    if value is None:
        raise ValueError("Upload a music file to begin.")
    if not isinstance(original_name, str) or not original_name.strip():
        raise ValueError("The upload is missing its original filename. Please select it again.")
    data = value if isinstance(value, gr.FileData) else gr.FileData.model_validate(value)
    path = os.path.realpath(data.path)
    if not os.path.isfile(path):
        raise ValueError("The uploaded file is no longer available. Please select it again.")
    return path, upload_wave_name(original_name)


def _validate_controls(course, level, bpm, temperature, top_p, drum_volume):
    if course not in COURSE_DENS:
        raise ValueError(f"Unsupported difficulty: {course}")
    if not 1 <= int(level) <= 10:
        raise ValueError("Level must be between 1 and 10 stars.")
    if bpm != 0 and not 30 <= float(bpm) <= 400:
        raise ValueError("Manual BPM must be between 30 and 400; use 0 for auto-detection.")
    if not 0.2 <= float(temperature) <= 1.2:
        raise ValueError("Temperature must be between 0.2 and 1.2.")
    if not 0.5 <= float(top_p) <= 1.0:
        raise ValueError("Top-p must be between 0.5 and 1.0.")
    if not 0 <= float(drum_volume) <= 1.5:
        raise ValueError("Taiko volume must be between 0% and 150%.")


def _file_data(path, *, name=None, mime_type=None):
    return gr.FileData(
        path=os.path.realpath(path),
        orig_name=name or os.path.basename(path),
        mime_type=mime_type,
    ).model_dump()


def _ura_notices():
    """What a user picking Ura has to be told, whichever model is selected."""
    corpus = {key: text for key, _, text in sc2_loader.KNOWN_LIMITATIONS}["ura"]
    return [
        {"when": "always", "text": corpus},
        {"when": "planner",
         "text": "The song planner was never trained on Ura, so its Oni plan "
                 "is used instead. Turn the AI planner off to avoid the "
                 "substitution."},
    ]


def capabilities():
    """Everything the interface needs to build its menus honestly.

    The BarScript entry is present only when its package resolves, and it carries
    the model card's limitations verbatim so the interface cannot paraphrase
    them into something milder.
    """
    models = []
    for name in available_models():
        group, label = MODEL_LABELS.get(name, (name, name))
        entry = {"value": name, "group": group, "label": label,
                 "experimental": is_sc2_model(name), "limitations": []}
        if is_sc2_model(name):
            entry["limitations"] = [
                {"section": section, "text": text}
                for _, section, text in sc2_loader.KNOWN_LIMITATIONS
            ]
            entry["notes"] = [
                "Song structure planning does not apply: BarScript takes no plan, "
                "so the Structure and AI planner switches are ignored and no "
                "#GOGO sections are written.",
                "BarScript needs a beat grid. Leave the beat grid on, or enter the "
                "song's BPM, so its bar lattice has a downbeat to sit on.",
                "Tempo changes are followed when the beat fit passes its quality "
                "gate: the estimated barlines become the bars, so the chart and "
                "its #BPMCHANGE lines track the song. If the fit fails the gate "
                "the bars are tiled from one BPM instead and a tempo change is "
                "not represented β€” the result says which happened.",
                "Time-signature changes are never followed. Nothing in this "
                "system estimates meter per bar from audio, so one estimated "
                "time signature is used for the whole song.",
            ]
        models.append(entry)
    return {
        "default_model": DEFAULT_MODEL,
        "models": models,
        "course_notices": {"ura": _ura_notices()},
    }


app = gr.Server(
    title="SoftChart",
    description="Conditional Taiko chart generation with synchronized audio preview.",
)
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")


@app.get("/", include_in_schema=False)
async def homepage():
    return FileResponse(
        STATIC_DIR / "index.html",
        headers={"Cache-Control": "no-cache"},
    )


@app.get("/health", include_in_schema=False)
async def health():
    return {"status": "ok", "device": DEVICE, "models_loaded": "gen" in _MODELS}


@app.get("/api/capabilities", include_in_schema=False)
async def api_capabilities():
    return capabilities()


@app.api(
    name="generate_chart",
    description="Generate a TJA chart, visual previews, and a synchronized taiko audio mix.",
    concurrency_limit=1,
    concurrency_id="generation-gpu",
    queue=True,
    stream_every=0.1,
)
def generate_chart(
    audio: gr.FileData,
    audio_name: str,
    course: str,
    level: int,
    bpm_override: float,
    auto_plan_on: bool,
    use_beat: bool,
    use_planner: bool,
    sampling: bool,
    temperature: float,
    top_p: float,
    drum_volume: float,
    model_choice: str = DEFAULT_MODEL,
) -> dict[str, object]:
    """Stream the actual inference stages to the custom frontend."""
    workdir = tempfile.mkdtemp(prefix="softchart-request-")
    try:
        audio_path, wave_name = _uploaded_file(audio, audio_name)
        _validate_controls(course, level, bpm_override, temperature, top_p, drum_volume)
        level = int(level)
        bpm_override = float(bpm_override)
        temperature = float(temperature)
        top_p = float(top_p)
        drum_volume = float(drum_volume)

        choices = available_models()
        if model_choice not in choices:
            raise ValueError(
                f"Unknown model {model_choice!r}. Choose one of: {', '.join(choices)}."
            )
        sc2_used = is_sc2_model(model_choice)
        yield _progress(
            "loading", 0.03, "Preparing",
            f"Loading SoftChart {model_choice}.",
        )
        if sc2_used:
            # BarScript lives in its own token space and cannot share a checkpoint
            # with the 1.x route; it borrows only the beat model, which it has
            # no equivalent of. include_planner is not honoured -- see below.
            sc2 = get_sc2()
            models = {"beat": sc2_beat_model()}
        else:
            sc2 = None
            models = get_models(model_choice, include_planner=use_planner)

        yield _progress(
            "audio", 0.11, "Listening",
            "Mapping rhythm, melody, and timbre.",
        )
        mel, wav = load_logmel(audio_path)

        yield _progress(
            "beat", 0.21, "Finding the beat",
            "Finding beats and tempo.",
        )
        # ``grid`` is the fit that PASSED its own quality gate (None otherwise).
        # ``beat_fit_raw`` keeps the last fit that was actually computed, gate
        # or no gate, so a failed gate can be reported with its numbers instead
        # of disappearing into a silent fallback.
        grid = dbs = beat_fit_raw = None
        if use_beat:
            grid = fit_grid_piecewise(models["beat"], mel, device=DEVICE)
            beat_fit_raw = grid
            if grid is not None:
                dbs = grid["downbeats"] if grid["ok"] else grid["db_peaks"]
                if not grid["ok"]:
                    grid = None
        if bpm_override > 0:
            bpm = bpm_override
            if use_beat and (grid is None or abs(grid["bpm"] - bpm) > 0.5):
                fixed_grid = fit_grid_fixed_bpm(models["beat"], mel, bpm, device=DEVICE)
                if fixed_grid is not None:
                    beat_fit_raw = fixed_grid
                grid = fixed_grid if fixed_grid is not None and fixed_grid["ok"] else None
                if grid is not None:
                    dbs = grid["downbeats"]
        elif grid is not None:
            bpm = float(grid["bpm"])
        elif dbs is not None and len(dbs) > 4:
            period = float(np.median(np.diff(dbs)))
            if period <= 0:
                raise RuntimeError("The beat model returned an invalid downbeat interval.")
            bpm = 240.0 / period
            while bpm >= 210:
                bpm /= 2.0
            while bpm < 70:
                bpm *= 2.0
        else:
            import librosa
            bpm = float(np.atleast_1d(librosa.beat.beat_track(y=wav, sr=SR)[0])[0])
        if not np.isfinite(bpm) or bpm <= 0:
            raise RuntimeError("Could not determine a reliable BPM. Enter it in Advanced settings.")
        if grid is None and abs(bpm - round(bpm)) < 0.06:
            bpm = float(round(bpm))

        yield _progress(
            "plan", 0.34, "Shaping the arc",
            "Planning density and climaxes.",
        )
        plan = None
        plan_source = "None"
        if sc2_used:
            # BarScript's decoder takes no plan token, so neither switch reaches it.
            # Reporting a plan that nothing consumed would be a lie in the UI.
            plan_source = "Not used β€” BarScript takes no plan"
        elif getattr(models["gen"], "_has_plan", False):
            if use_planner:
                asked, substituted = planner_course(course)
                plan = learned_plan(models["plan"], mel, course, bpm, dbs)
                plan_source = (f"AI planner ({asked.capitalize()} plan reused "
                               f"for {course.capitalize()})" if substituted
                               else "AI planner")
            elif auto_plan_on:
                plan = auto_plan(mel, bpm, dbs)
                plan_source = "Energy heuristic"

        yield _progress(
            "generate", 0.47, "Writing the chart",
            "Writing playable Taiko patterns.",
        )
        title = os.path.splitext(wave_name)[0]
        extra_metrics = {}
        slot_used = grid is not None
        if sc2_used:
            generated, tja, extra_metrics = generate_sc2(
                sc2, mel, grid, beat_fit_raw, dbs, bpm, course, level,
                wave_name, title,
                sampling=sampling, temperature=temperature, top_p=top_p,
            )
            slot_used = True
        elif slot_used:
            generated = generate_song_slot(
                models["slot"], mel, grid, course, level=level,
                density_bucket=COURSE_DENS[course], greedy=not sampling, seed=0,
                temperature=temperature, top_p=top_p, device=DEVICE, plan=plan,
            )
            tja = normalize_course_header(write_tja_slots(
                generated, grid, title, course, level, wave_name, plan=plan,
            ), course)
        else:
            generated = generate_song(
                models["gen"], mel, course, level=level,
                density_bucket=COURSE_DENS[course], greedy=not sampling,
                temperature=temperature, top_p=top_p, seed=0,
                device=DEVICE, plan=plan,
            )
            generated = snap_chart(generated, bpm)
            tja = write_tja(
                generated, bpm, title, course, level, wave_name, dbs,
                grid_fit=grid, plan=plan,
            )

        yield _progress(
            "export", 0.81, "Rendering",
            "Building the TJA and previews.",
        )
        tja_path = output_tja_path(wave_name, course, workdir)
        with open(tja_path, "w", encoding="utf-8") as output_file:
            output_file.write(tja)
        chart_image_path = os.path.join(workdir, "chart.png")
        plan_image_path = os.path.join(workdir, "song-plan.png")
        render_tja_image(tja, out_path=chart_image_path)
        render_audio_plan(mel, title, course, plan=plan, out_path=plan_image_path)

        yield _progress(
            "mix", 0.91, "Mixing",
            "Mixing Taiko with your track.",
        )
        stem = Path(tja_path).stem
        preview_path = os.path.join(workdir, f"{stem}_taiko-preview.wav")
        mix_stats = synthesize_taiko_preview(
            audio_path,
            tja,
            preview_path,
            drum_gain=drum_volume,
            sample_rate=44100,
        )

        grid_rms = round(float(grid["rms_ms"]), 1) if grid is not None else None
        metrics = {
            "bpm": round(float(bpm), 1),
            "notes": len(generated["hits"]),
            "spans": len(generated["spans"]),
            "timing": "slot-exact" if slot_used else "time-quantized",
            "grid_rms_ms": grid_rms,
            "preview_hits": int(mix_stats["rendered_hit_count"]),
            "model": model_choice,
            "difficulty": f"{course.capitalize()} (TJA {TJA_COURSE_HEADER[course]})",
            "plan_source": plan_source,
            **extra_metrics,
        }
        yield {
            "kind": "complete",
            "stage": "complete",
            "progress": 1.0,
            "title": "Ready",
            "detail": "Play it or download it.",
            "metrics": metrics,
            "files": {
                "tja": _file_data(tja_path, mime_type="text/plain"),
                "audio": _file_data(preview_path, mime_type="audio/wav"),
                "chart_image": _file_data(chart_image_path, mime_type="image/png"),
                "plan_image": _file_data(plan_image_path, mime_type="image/png"),
            },
        }
    except Exception as exc:
        LOGGER.exception("SoftChart generation failed")
        detail = (str(exc) if isinstance(exc, (ValueError, RuntimeError))
                  else "The server could not complete this chart. Please try again shortly.")
        yield {
            "kind": "error",
            "stage": "error",
            "progress": 0.0,
            "title": "Generation did not complete",
            "detail": detail,
        }
    finally:
        shutil.rmtree(workdir, ignore_errors=True)


if __name__ == "__main__":
    app.launch(max_file_size="200mb")