Spaces:
Sleeping
Sleeping
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")
|