diff --git a/REPORT.md b/REPORT.md index e8afa7d..4dafa59 100644 --- a/REPORT.md +++ b/REPORT.md @@ -62,23 +62,70 @@ Ten sam cel (predykcja następnego znaku), spektrum mechanizmów — od twardego - **GPT** wygrywa ppl zmiennym, długim kontekstem (uwaga), kosztem ~800K param. Prawdziwe osie spektrum to **pamięć i generalizacja**, nie sam ppl. -## 5. Co wiemy / czego nie - -**Udowodnione (z liczbą):** mały char-LM uczy się struktury muzyki; czyszczenie danych obniża perplexity (3,88 → 3,80); mechanizm szwu bezstratny; ensemble bije pojedyncze modele; **niezależne maluchy mają wspólną geometrię** (CKA, po audycie); zmierzone **spektrum n-gram → NPLM → transformer** (ppl vs pamięć vs generalizacja). - -**Jeszcze nie:** że łączenie reprezentacji **bije** ensemble — to wymaga ekspertów **komplementarnych** (różne, nakładające się domeny) i/lub wymuszonego wspólnego kontraktu; routing. - -## 6. Następne kroki +## 5. Łuk trzeci: sędzia — mierzenie „czy to w ogóle muzyka" (E-JUDGE) + +PPL ma dwie ślepe plamy: nie odpowiada, czy generacja *działa jako muzyka*, i nie porównuje +modeli z różnych korpusów. Zbudowaliśmy więc niezależnego sędziego: klasyfikator `JudgeGPT` +(trunk GPT verbatim + mean-pool + 3 głowy; 55 tys. melodii z thesession.org, split po +`tune_id`), który patrząc TYLKO na ciało melodii (nagłówki `M:`/`K:` wycinane — nie można +oszukiwać promptem) przewiduje **meter / mode / type**. Val acc: meter 0.92, mode 0.72, +type 0.82 — z pomyłkami muzycznie prawidłowymi (hornpipe↔reel to para najtrudniejsza nawet +dla ludzi; Bmin↔D to ten sam zbiór dźwięków — sufit informacyjny). + +**Benchmark domenowy.** Każdego eksperta oceniamy wyłącznie na melodiach z jego własnej +domeny (pula walidacyjna ograniczona filtrem type+metrum z `domains.json`). Wynik (score) to +średnie prawdopodobieństwo, jakie sędzia przyznaje prawdziwym etykietom; każda porażka modelu +liczy się jako zero (brak zniekształcenia przez „przeżywających"), a punktem odniesienia +(sufitem) jest ocena przez tego samego sędziego **prawdziwych** melodii tej domeny. Wyniki: +jigi osiągają **80–86% sufitu**, walc 86% (sufit miękki — sędzia sam jest mało pewny na 3/4), +reele 69–72%. Ranking wyznaczony przez sędziego pokrywa się z rankingiem PPL — dwie +niezależne miary, jeden werdykt. Odnotujmy uczciwie: pierwsza wersja benchmarku, oceniająca +wszystkich ekspertów na jednej wspólnej puli melodii, pokazywała dokładnie odwrotnie +(najlepiej wypadały reele). To była wada pomiaru, nie właściwość modeli: niezbalansowane +klasy sprawiały, że ekspert 4/4 „wygrywał" tam, gdzie prawda o metrum była po prostu +najczęstsza. Dopiero pule domenowe dały uczciwy obraz — kolejny pomiar, który obalił nasze +pierwsze odczytanie. + +**Macierz OOD** (ekspert × domena celu; prawda = to, o co prosimy): +- Diagonala dominuje wszędzie — **specjalizacja ekspertów jest realna**. +- **Nagłówki nie sterują poza domeną**: jig z promptem `M:4/4` generuje 6/8 z pewnością + 0.94 (home-bias 0.9). Sterowanie promptem jest martwe. +- **Kontekst steruje w połowie**: z prawdziwym prefiksem melodii home-bias spada do ~0.5, + score rośnie do 40–63% sufitu. Modele czytają styl z nut, nie z nagłówków — w domenie też + (mode z `K:` słabo, z prefiksu lepiej). +- **Asymetria**: najlepsze komórki OOD to reele→jigi (≈60% sufitu, home-bias type 0.2) — + eksperci 4/4 są elastyczni, 6/8 to twardy nawyk. +- **Elastyczność siedzi w małych modelach**: e32 mają najniższy home-bias wszędzie + (najsłabiej „zapatrzone" w swoją domenę) przy najsłabszej diagonali; duże modele są + najsztywniejsze. Nowa hipoteza pod E1: **komplementarność może wolić małe, gibkie modele + mimo gorszego PPL**. +- Bach (chorały) jest uczciwie wyłączony: bez wspólnego słownika nie da się go nawet + zakodować w irlandzkich promptach — wspólny kontrakt słownikowy to warunek wstępny + dalszych testów. + +Ograniczenia: sędzia ma fizyczny sufit na mode (tonacje względne są niedróżnialne z treści +dźwiękowej); macierz OOD policzona na n=30/komórkę — wersja pod paper wymaga pełnego +przebiegu i sufitu z całej puli. + +## 6. Co wiemy / czego nie + +**Udowodnione (z liczbą):** mały char-LM uczy się struktury muzyki; czyszczenie danych obniża perplexity (3,88 → 3,80); mechanizm szwu bezstratny; ensemble bije pojedyncze modele; **niezależne maluchy mają wspólną geometrię** (CKA, po audycie); zmierzone **spektrum n-gram → NPLM → transformer** (ppl vs pamięć vs generalizacja); **sędzia niezależny od PPL potwierdza ranking ekspertów** i mierzy nową oś (specjalizacja: jigi 80–86% sufitu, reele 69–72%); **nagłówki nie sterują ekspertem poza domeną, kontekst melodyczny tak (w połowie)**; **elastyczność OOD rośnie, gdy model jest mniejszy** (hipoteza pod E1, do potwierdzenia). + +**Jeszcze nie:** że łączenie reprezentacji **bije** ensemble — to wymaga ekspertów **komplementarnych** (różne, nakładające się domeny) i/lub wymuszonego wspólnego kontraktu; routing; pełna (n=100) macierz OOD z sufitem z całej puli; wspólny słownik (warunek testów międzykorpusem, m.in. dla Bacha). + +## 7. Następne kroki 1. Pre-check: czy rozbieżność (wariancja) między ekspertami koreluje z błędem — bramka przed routingiem. 2. Wymuszony wspólny kontrakt (zamrożony front + głowa) na stylach w **tym samym metrum** (różne metrum = model oszukuje po nagłówku — pułapka pomiaru) + ekspertach **komplementarnych**. -3. Bogatszy, nieliniowy łącznik — dopiero jeśli (2) pokaże sygnał. +3. **Wspólny słownik przy trenowaniu ekspertów** — usuwa klasę artefaktów OOD (dziś: pominięcia słownikowe, wyłączony Bach) i jest wstępem do stitcha między ekspertami. +4. E-JUDGE ensemble/stitch: te same melodie, ten sam sędzia — czy metoda kompozycji bije pojedynczych ekspertów w oś stylu, nie tylko w PPL. +5. Bogatszy, nieliniowy łącznik — dopiero jeśli (2) pokaże sygnał. - Osobny kierunek: wiele modeli „grające razem" (polifonia, synchronizacja). -## 7. Po co to +## 8. Po co to Dwa cele: (a) **budować know-how zespołu Slayer** — tani, jawny, reprodukowalny poligon; (b) **de-ryzykować metody** pod docelowy model budowlany (klasyfikacja / tworzenie / rozumienie `.ifc`), gdzie poprawność jest sprawdzalna kodem (walidator = weryfikowalna nagroda). ## Metoda w akcji (dlaczego to wiarygodne) -W trakcie tych prac **pomiar dwukrotnie obalił wewnętrzną hipotezę zespołu** i zostało to przyjęte, nie naciągnięte: (1) pierwszy E1 „bił ensemble" — okazał się artefaktem zbioru testowego; (2) E_CKA przeszedł **adwersarialny audyt własnego pomiaru**, który wycofał jedną z dwóch metryk (skażony baseline). To jest sedno tego labu: *najpierw mierz, potem twierdź — i audytuj własny pomiar.* +W trakcie tych prac **pomiar trzykrotnie obalił wewnętrzne przekonania zespołu** i zostało to przyjęte, nie naciągnięte: (1) pierwszy E1 „bił ensemble" — okazał się artefaktem zbioru testowego; (2) E_CKA przeszedł **adwersarialny audyt własnego pomiaru**, który wycofał jedną z dwóch metryk (skażony baseline); (3) pierwszy benchmark sędziego (wspólna pula dla wszystkich ekspertów) pokazał ranking odwrotny niż PPL — okazał się artefaktem niezbalansowanych klas; naprawa: pule domenowe. To jest sedno tego labu: *najpierw mierz, potem twierdź — i audytuj własny pomiar.* --- *Pełne rozumowanie (teza ↔ antyteza → synteza, audyt, weryfikacja źródeł) i kod: `music-experts/docs/` oraz `music-experts/src/`.* diff --git a/music-experts/.gitignore b/music-experts/.gitignore index 106d345..73b88eb 100644 --- a/music-experts/.gitignore +++ b/music-experts/.gitignore @@ -1,11 +1,12 @@ # cache __pycache__/ *.pyc - +.venv # Repo = kod + wagi + docs. Surowe dane i OUTPUTY (nagrania, korpusy, raporty robocze) # są regenerowalne, więc NIE idą. Idą tylko wytrenowane MODELE (wagi). data/* -!data/models +out/* +data/models data/models/* !data/models/*.pt data/models/gpt_ckpt_v1_dirty.pt diff --git a/music-experts/README.md b/music-experts/README.md index 41e7f66..86761ce 100644 --- a/music-experts/README.md +++ b/music-experts/README.md @@ -40,10 +40,45 @@ context 128 chars, character-level. Trained on CPU in minutes. - **Duet** (`src/compose/duet.py`) — two experts layered (piano + violin, simultaneous): multi-track, not model-level fusion. - **Next — E1:** representation-level stitch (the actual hypothesis, meant to beat these baselines). +## Judge — automatic generation benchmark + +Perplexity measures text fit, not whether a generated tune *works as music* (meter? key? dance +type?). The judge is an independent classifier that scores generated melodies like a human +would: seeing only the tune **body** (the `M:`/`K:` headers are stripped — the judge cannot +cheat off the prompt), it predicts **meter**, **mode** and **type** (jig/reel/hornpipe...). +Built from the same building blocks as the experts: `JudgeGPT` = `core/gpt.py` trunk (verbatim) ++ mean-pool + 3 linear heads, pure PyTorch; 55k tunes, leak-free split by `tune_id`. +`judge_v2` val acc: meter **0.92**, mode **0.72**, type **0.82** — with musically honest +confusions (hornpipe↔reel, Bmin↔D). + +Two benchmarks built on top of it: + +- **In-domain** (`benchmark_judge.py`): each expert is scored only on tunes from its own + domain (`data/models/domains.json`), failures count as 0 (no survivor bias), with a + real-melody ceiling for reference. Leaderboard: jigs hold **0.80–0.86** of ceiling, + waltz 0.86 (soft ceiling), reels 0.69–0.72. The judge's ranking agrees with PPL — two + independent metrics, one verdict. +- **OOD transfer matrix** (`benchmark_ood.py`): expert × target domain. Diagonal always + wins (specialization is real), but **headers don't steer off-domain** (home-bias 0.9 — + a jig expert told `M:4/4` still emits 6/8), while **melody context half-steers** (home-bias + ~0.5). Reel experts are the most flexible donors; small (e32) models bend easiest, + big ones are most rigid. + +```bash +./run_benchmarks.sh # in-domain sweep, all checkpoints +python src/tools/train_judge.py --iters 5000 --batch 64 --out data/models/judge_v2.pt +python src/tools/benchmark_judge.py --model data/models/jig_ckpt.pt --judge data/models/judge_v2.pt +python src/tools/benchmark_ood.py # expert × domain transfer matrix +``` + +Full WHY / HOW / WHAT, results and limitations (in Polish): +[docs/Badania/2026-09-05_posttraining-reverse-kl/Judge-Sedzia-Generacji.md](docs/Badania/2026-09-05_posttraining-reverse-kl/Judge-Sedzia-Generacji.md) + ## Pipeline (`src/`) `prepare_data.py` / `prepare_bach.py` (build ABC corpus) → `gpt.py` (architecture) → `train_gpt.py` (train; optional shared vocab) → `make_midi.py` / `gen_samples.py` (generate + render) → -`e0_stitch.py` / `fuse.py` / `duet.py` (composition) · `ngram_model.py` (baseline) · `abc_to_midi.py` (render). +`e0_stitch.py` / `fuse.py` / `duet.py` (composition) · `ngram_model.py` (baseline) · `abc_to_midi.py` (render) · +`train_judge.py` / `judge_tunes.py` / `benchmark_judge.py` / `benchmark_ood.py` (judge — generation benchmark). ## Usage ```bash diff --git a/music-experts/data/models/bach_ckpt.pt b/music-experts/data/models/bach_ckpt.pt deleted file mode 100644 index 316e49d..0000000 Binary files a/music-experts/data/models/bach_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/domains.json b/music-experts/data/models/domains.json new file mode 100644 index 0000000..0b14c19 --- /dev/null +++ b/music-experts/data/models/domains.json @@ -0,0 +1,31 @@ +{ + "_author": "Adam Skrodzki", + "_comment": "Domena każdego checkpointu (co przeszło do korpusu przy prepare_data.py). Domena = filtr melodii z tunes.csv: 'type' to substring typu, 'meter' równość. null = model spoza zakresu benchmarku (brak dopasowanych melodii).", + "bach_ckpt.pt": {"type": null, "meter": null}, + "jig_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_v2_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e32_s1_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e32_s2_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e32_s3_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e64_s3_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e128_s3_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_e192_s3_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_l1_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_l2_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_s1_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "jig_s2_ckpt.pt": {"type": "jig", "meter": "6/8"}, + "reel_ckpt.pt": {"type": "reel", "meter": "4/4"}, + "reel_sv_ckpt.pt": {"type": "reel", "meter": "4/4"}, + "waltz_ckpt.pt": {"type": "waltz", "meter": "3/4"}, + "universalist_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna: model trenowany na calym korpusie; type wplywa tylko na gwiazdke diag i glowe hb"}, + "jig_sh_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "teacher RKL, wspolny slownik (VOCAB_FROM=universalist)"}, + "reel_sh_ckpt.pt": {"type": "reel", "meter": "4/4", "_note": "teacher RKL, wspolny slownik (VOCAB_FROM=universalist)"}, + "waltz_sh_ckpt.pt": {"type": "waltz", "meter": "3/4", "_note": "teacher RKL, wspolny slownik (VOCAB_FROM=universalist)"}, + "student_rkl_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna (student po RKL); j.w."}, + "student_junk_rkl_sc_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna (student po RKL); j.w."}, + "student_rkl_last_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna (student po RKL, last); j.w."}, + "student_rkl_sc_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna (student po RKL cont+scratch); j.w."}, + "student_rkl_sc_last_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "domena arbitralna (student po RKL cont+scratch, last); j.w."}, + "student_jigstart_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "ablation: start z jig_sh (bez widzenia reel/waltz), RKL cont+scratch; domena arbitralna"}, + "student_jigstart_last_ckpt.pt": {"type": "jig", "meter": "6/8", "_note": "j.w., last"} +} diff --git a/music-experts/data/models/jig_ckpt.pt b/music-experts/data/models/jig_ckpt.pt deleted file mode 100644 index 6a65e8a..0000000 Binary files a/music-experts/data/models/jig_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e128_s3_ckpt.pt b/music-experts/data/models/jig_e128_s3_ckpt.pt deleted file mode 100644 index 32f756e..0000000 Binary files a/music-experts/data/models/jig_e128_s3_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e192_s3_ckpt.pt b/music-experts/data/models/jig_e192_s3_ckpt.pt deleted file mode 100644 index e49e623..0000000 Binary files a/music-experts/data/models/jig_e192_s3_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e32_s1_ckpt.pt b/music-experts/data/models/jig_e32_s1_ckpt.pt deleted file mode 100644 index eaa5e9f..0000000 Binary files a/music-experts/data/models/jig_e32_s1_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e32_s2_ckpt.pt b/music-experts/data/models/jig_e32_s2_ckpt.pt deleted file mode 100644 index f823108..0000000 Binary files a/music-experts/data/models/jig_e32_s2_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e32_s3_ckpt.pt b/music-experts/data/models/jig_e32_s3_ckpt.pt deleted file mode 100644 index 7b76c32..0000000 Binary files a/music-experts/data/models/jig_e32_s3_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_e64_s3_ckpt.pt b/music-experts/data/models/jig_e64_s3_ckpt.pt deleted file mode 100644 index a70a65a..0000000 Binary files a/music-experts/data/models/jig_e64_s3_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_l1_ckpt.pt b/music-experts/data/models/jig_l1_ckpt.pt deleted file mode 100644 index 72aecea..0000000 Binary files a/music-experts/data/models/jig_l1_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_l2_ckpt.pt b/music-experts/data/models/jig_l2_ckpt.pt deleted file mode 100644 index 40b8a01..0000000 Binary files a/music-experts/data/models/jig_l2_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_s1_ckpt.pt b/music-experts/data/models/jig_s1_ckpt.pt deleted file mode 100644 index d85eedb..0000000 Binary files a/music-experts/data/models/jig_s1_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_s2_ckpt.pt b/music-experts/data/models/jig_s2_ckpt.pt deleted file mode 100644 index 5573690..0000000 Binary files a/music-experts/data/models/jig_s2_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/jig_v2_ckpt.pt b/music-experts/data/models/jig_v2_ckpt.pt deleted file mode 100644 index 4959fe9..0000000 Binary files a/music-experts/data/models/jig_v2_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/reel_ckpt.pt b/music-experts/data/models/reel_ckpt.pt deleted file mode 100644 index 86b0550..0000000 Binary files a/music-experts/data/models/reel_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/reel_sv_ckpt.pt b/music-experts/data/models/reel_sv_ckpt.pt deleted file mode 100644 index 8f1ae2e..0000000 Binary files a/music-experts/data/models/reel_sv_ckpt.pt and /dev/null differ diff --git a/music-experts/data/models/waltz_ckpt.pt b/music-experts/data/models/waltz_ckpt.pt deleted file mode 100644 index 6eb3fc6..0000000 Binary files a/music-experts/data/models/waltz_ckpt.pt and /dev/null differ diff --git a/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Judge-Sedzia-Generacji.md b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Judge-Sedzia-Generacji.md new file mode 100644 index 0000000..9905b2c --- /dev/null +++ b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Judge-Sedzia-Generacji.md @@ -0,0 +1,248 @@ +--- +type: koncepcja +title: "Judge — automatyczny sędzia jakości generacji (benchmark)" +status: aktywne +data: 2026-09-04 +author: Adam Skrodzki +related: "[[README]] · src/core/judge.py · src/tools/train_judge.py · src/tools/judge_tunes.py" +--- + +# Judge — automatyczny sędzia jakości generacji + +## WHY — po co to jest + +### Problem: perplexity nie mierzy jakości muzyki + +Nasza jedyna dotychczasowa miara jakości eksperta to **perplexity na held-oucie** — czyli +pytanie „jak dobrze model przewiduje następny znak tekstu, na którym się nie trenował". To +uczciwa miara *dopasowania do korpusu*, ale ma dwie ślepe plamy: + +1. **PPL nie odpowiada na pytanie „czy to jest muzyka?"** — model może mieć niską PPL i + generować ciągi znaków ABC, które *wyglądają* jak nuty, ale nie trzymają metrum, tonacji + ani rytmu typowego dla tańca. Do tej pory to sprawdzaliśmy **uchem** — qualitatywnie, + niereprodukowalnie, bez liczby do tabeli. +2. **PPL nie porównuje modeli trenowanych na różnych korpusach** (jak Bach vs jig — + dopisek w README: „ppl nie jest wprost porównywalne"). + +Potrzebna jest **niezależna, automatyczna, reprodukowalna miara treści muzycznej** — sędzia, +który spojrzy na wygenerowaną melodię i odpowie na pytania, na które odpowiedziałby człowiek +znający irlandzki trad: *jakie to metrum? jaka tonacja? do jakiego tańca to pasuje?* + +### Idea: klasyfikator jako sędzia (classifier-as-judge) + +Trenujemy klasyfikator na **prawdziwych** melodiach, które mają +etykiety: `meter` (7 klas), `mode` (23 klasy), `type` (12 klas). Potem pokazujemy mu +**wygenerowane** melodie i sprawdzamy zgodność: + +- ekspert `jig_ckpt` generuje z promptem `M:6/8 K:D` → sędzia powinien powiedzieć + `6/8 + D + jig`. **Zgoda = model trzyma styl.** +- sędzia mówi `4/4 + reel` → **dryf stylu** — model „wyszedł poza swoją specjalizację". + +To daje liczbę (procent zgodności), którą można kłaść do tabeli obok PPL — i którą można +porównywać *między eksperymentami kompozycji* (E0/E1/ensemble: który spójnie łączy ekspertów, +a który produkuje mush?). Ostatecznie to pierwszy krok do odpowiedzi na pytanie z KMC1: +**czy kompozycja małych modeli daje w ogóle mierzalny zysk?** + +### Czemu sędzia widzi tylko ciało melodii + +Generowane melodie zawsze mają nagłówek `X:1 M:6/8 K:D` — bo to jest prompt. Gdyby sędzia +widział nagłówek, nauczyłby się czytać odpowiedź z promptu („M:6/8 → meter=6/8", gotowe, +accuracy 100%) i niczego nie mierzył. Dlatego featury budujemy **wyłącznie z ciała melodii** +(nut), a nagłówki `X:/T:/C:/M:/K:/N:/%` są wycinane przed featuryzacją. Sędzia musi wnioskować +metrum i tonację z **rytmu i rozkładu dźwięków** — tak jak człowiek. + +## WHAT — co to jest + +### Budowa: `JudgeGPT` (`src/core/judge.py`) + +Klocki **te same co eksperci** — verbatim reuse `core/gpt.py` + +``` +ciało melodii (znaki, do 256) ─┐ + ├─ trunk GPT (verbatim: tok_emb+pos_emb → 4×Block → ln_f) +maska paddingu ─────────────────┘ │ + mean-pool po pozycjach (tylko prawdziwe znaki) + │ + ┌─────────────┼─────────────┐ + głowa meter głowa mode głowa type (nn.Linear, ~n_klas wyjść) +``` + +Kluczowe decyzje projektowe: + +- **CausalAttention wystarcza.** Nie pisaliśmy bidirectional attention. Padding jest na + KOŃCU sekwencji, więc tokeny prawdziwe nigdy nie widzą paddingu (maska dolnotrójkątna to + gwarantuje), a maska jest potrzebna tylko do poprawnego mean-poolingu. Reprezentacja po + poolingu zagregowana po całej melodii — kierunkowość uwagi przestaje mieć znaczenie. +- **Trunk = cały GPT**, razem z nieużywaną głową LM (weight tying z embeddingiem sprawia, że + nie kosztuje dodatkowej pamięci znaczącej). To świadomy wybór: nie „wyciągamy" bloków + ręcznie, tylko używamy klasy `GPT` taką, jaka jest — mniej kodu, zero rozjazdu z eksper- + tami, i otwarta droga do `--init-from` (transfer z pretrenowanego eksperta). +- **Okna ciała 256 znaków.** Ciała mają 40–700 znaków. Trening: losowe okno na przykład na + każdy krok (augmentacja — przez czas treningu model widzi całe ciało, nie tylko początek). + Eval: wszystkie okna niepokrywające, softmax uśredniony po oknach → jedna predykcja na + melodię (pełne, deterministyczne pokrycie). + +### Dane i split (`data/tunes.csv`) + +- Czyszczenie **współdzielone z prepare_data.py** (wyciągnięte do `src/core/abc_corpus.py`): + usunięcie akordów `"..."`, normalizacja tonacji (`Edorian` → `Edor`), filtr długości ciała + 40–700 znaków. Sędzia widzi dokładnie tę samą dystrybucję, którą widział generator. +- **Split 90/10 po `tune_id`** — to ważne: jedna melodia ma często kilka settingów (rows w + csv) niemal identycznych. Losowy split po wierszach przeciekałby (copypasta w val), + split po `tune_id` gwarantuje, że val to melodie, których sędzia nigdy nie widział. + +### Zastrzeżenie: split chroni sędziego, ale nie generatory + +Split po `tune_id` gwarantuje czystość walidacji **sędziego**. Generatory (eksperci, +universalista, nauczyciele `_sh`) są trenowane na korpusach z `prepare_data.py`, które +biorą CAŁY `tunes.csv` — w tym melodie z val splitu sędziego. Sędzia nie wycieka do +treningu generatorów (żaden generator nie widzi jego wyjść ani wag), ale generator mógł +zapamiętać surowy tekst melodii walidacyjnych. Konsekwencje: + +- **Porównania międzymodelami są uczciwe** — wyciek jest wspólny dla wszystkich + benchmarkowanych generatorów (eksperci, universalista, studenci RKL). +- **Wartości absolutne są optymistyczne** — continuation jest najbardziej narażony + (prompt = ćwiartka ciała melodii walidacyjnej, którą generator mógł zapamiętać), + scratch najmniej (same nagłówki). Sufit (`ref`) jest niewrażliwy. +- **Clean-room wariant** (dla paperu): `prepare_data.py` z wykluczeniem val `tune_id` + splitu sędziego + przetrenowanie wszystkiego. + +### Wyniki + +| głowa | val acc (judge_v2) | baseline (klasa większościowa) | interpretacja | +|---|---|---|---| +| meter | **0.922** | ≈ 0.5 (4/4) | mocno ponad baseline | +| mode | **0.717** | ≈ 0.35 (D) | blisko sufitu informacyjnego (patrz niżej) | +| type | **0.816** | ≈ 0.35 (reel) | mocno ponad baseline | + +(pierwsza wersja sędziego, 2000 kroków, dawała 0.891/0.688/0.778 — judge_v2 to 5000 kroków, +batch 64; kanoniczny checkpoint: `data/models/judge_v2.pt`) + +**Top pomyłki są muzycznie prawidłowe** — i to jest najważniejszy wynik walidacyjny: + +- `hornpipe→reel` (185), `reel→hornpipe` (82): oba to 4/4, różnią się tylko **feelingiem** + rytmicznego puntowania. Nawet ludzie się spierają. +- `3/4→6/8`, `9/8→12/8`: grupowanie ósemek zamiast samego liczenia — podobieństwo rytmiczne. +- `Bmin→D`, `Edor→Emin`, `Amix→Ador`: **to ten sam zbiór dźwięków** (tonacja względna / + tryby o tych samych znakach). Z samego rozkładu nut te pary są *nieodróżnialne wprost* — + rozstrzyga je tylko waga kadencji. ~0.7 na mode jest więc blisko sufitu tego, co można + wywnioskować z zawartości dźwiękowej. + +Sędzia myli się tam, gdzie muzycznie *powinien* się mylić — czyli mierzy strukturę muzyczną, +a nie artefakty formatu. Val acc ≈ train acc: brak overfittingu. + +## HOW — jak tego używać + +### Trening + +```bash +python src/tools/train_judge.py +# warianty: +python src/tools/train_judge.py --init-from data/models/jig_ckpt.pt # trunk od eksperta (transfer learning — eksperyment do paperu) +python src/tools/train_judge.py --iters 5000 --lr 3e-4 # dłuższy trening +``` + +Pełna lista flag: `--csv --out --seed --iters --batch --lr --block --eval-every +--train-eval-n`. Co 100 kroków logowane są **train acc** (stała próbka 5000 przykładów) i +**val acc** (cały val); gap między nimi to sygnał overfittingu. Batch = 32 melodie/krok; +próbkowanie losowe bez epoch — ten sam przykład wraca średnio co `len(train)/batch` ≈ 1541 +kroków, za każdym razem z nowym losowym oknem. Model zapisuje best-val do +`data/models/judge_v2.pt` (domyślny `--out`). + +### Ocena wygenerowanych melodii + +```bash +python src/tools/judge_tunes.py out/ # pliki .abc i/lub katalogi (rekurencyjnie) +python src/tools/judge_tunes.py out/ --json # postać maszynowa (benchmark scripts) +``` + +Wynik: na każdą melodię predykcja meter/mode/type + pewność (softmax). Melodię z jednego +pliku wielomelodiowego skrypt tnie po liniach `X:` (ta sama konwencja co `first_tune` w +generatorach). + +### Interpretacja: zgodność z ekspertem + +Dla eksperta `E` generującego z promptu `M:m K:k` sprawdzamy: +- `judge.meter == m` (metryka z promptu), +- `judge.mode == k` (tonacja z promptu), +- `judge.type == typ korpusu E` (jig→jig, reel→reel...). + +**Procent zgodności = benchmark jakości generacji E.** Niska zgodność przy dobrym PPL to +sygnał, że model nauczył się powierzchni tekstu, a nie struktury muzycznej. Ten sam miernik +możemy przyłożyć do ensemble'u i stitchy — porównanie zgodności między metodami kompozycji +to główny przewidywany eksperyment (E-JUDGE). + +## Benchmark ekspertów (E-JUDGE): domena i OOD + +### Benchmark domenowy — `benchmark_judge.py` + +Eksperta oceniamy TYLKO na jego własnym terenie. Pula = val split ograniczony do **domeny +checkpointu** (`data/models/domains.json`: substring typu + metrum, jak w `prepare_data.py`). +Dwie próby: **scratch** (tylko nagłówki `M:`/`K:`, model generuje od zera; type pominięty — +model nie ma jak go znać) i **continuation** (ćwiartka prawdziwego ciała w promptcie; sędzia +ocenia całość). Score = średnie P sędziego przy prawdziwych klasach, mianownik = WSZYSTKIE +próbki: porażka modelu w domenie = 0 (bez survivor bias). Referencja = sędzia na prawdziwych +melodiach tej domeny (sufit). Ten sam seed = identyczne melodie dla każdego modelu. + +**Wyniki (n=100, judge_v2, score/sufit):** + +| rodzina | continuation | stosunek do sufitu | scratch mode P | wniosek | +|---|---|---|---|---| +| jigi (11 wariantów) | 0.61–0.66 / 0.77 | **0.80–0.86** | 0.15–0.27 | najsilniejsza rodzina; `jig_l1` najlepszy | +| reele (2) | 0.53–0.55 / 0.76 | 0.69–0.72 | 0.23–0.27 | najsłabsi u siebie | +| walc | 0.48 / 0.56 | 0.86 (miękki sufit) | 0.17 | type = 1.00 sufitu, ale sufit niski | + +Kluczowe obserwacje: +- **Ranking sędziego zgadza się z rankingiem PPL** (jig 3.80 najlepszy) — dwie niezależne + miary, jeden werdykt. Pierwsza wersja benchmarku (krzyżowa, bez domen) dawała odwrotność — + artefakt bazowych częstości klas; przejście na pule domenowe to naprawiło. +- **Mode z nagłówka słabo, z kontekstu lepiej**: w scratch sędzia przypisuje prawdziwej + tonacji tylko ~0.2; w continuation ~0.35 (prefiks niesie materiał dźwiękowy). Modele + czytają tonację z nut, nie z `K:` — konkretna, sprawdzalna teza o tym, czego nauczył się + char-LM. +- **Zero śmieci w domenie**: wszystkie pominięcia to niedopasowania słownika (te same 7 wierszy + dla każdego jiga — deterministyczny seed), nie generacje-garbage. + +### Macierz OOD — `benchmark_ood.py` + +Wiersze = eksperci, kolumny = domeny celu. Prawda = to, o co *prosimy* (etykiety melodii +celu — model jest proszony, więc nie ma wymówki). Komórka: score/ref, **hb = home-bias** +(P sędzia przypisze KLASIE DOMOWEJ modelu: meter, a w continuation też type) i pokrycie. +hb rozróżnia dwa tryby porażki: wysoki hb = model sztywny, ignoruje prompt i wraca do +domeny; niski hb + niski score = papka. Bach wyłączony (`type: null`) — bez wspólnego +słownika nie da się uczciwie mierzyć (patrz „następny krok"). + +**Wyniki (n=30, judge_v2; diagonala = domena własna):** + +- **Diagonala dominuje wszędzie** — specjalizacja ekspertów jest realna, przeżyła pierwszy + adwersarialny test. +- **Nagłówki nie sterują poza domeną**: OOD scratch ma hb 0.85–0.97 — jig z promptem `M:4/4` + generuje 6/8 z pewnością 0.94. Sterowanie promptem jest martwe. +- **Kontekst steruje w połowie**: OOD continuation spada hb do ~0.40–0.55, score rośnie do + 40–63% sufitu. Wniosek projektowy: w kompozycji steruj kontekstem (stitch), nie nagłówkami. +- **Asymetria reele→jigi**: najlepsze komórki OOD w całej macierzy (0.48–0.52, hb type + zaledwie 0.20–0.24) — eksperci 4/4 są elastyczni, 6/8 to trudniejszy nawyk do zdjęcia. + Kandydaci na „dawców" w modelach mieszanych. +- **Elastyczność siedzi w małych modelach**: e32 mają najniższy hb wszędzie (0.62–0.79) + przy najsłabszej diagonali; duże modele hb 0.90–0.97. Nowa, testowalna hipoteza pod E1: + **komplementarność może wolić małe, „gibkie" modele mimo gorszego PPL.** +- Walc na diagonali osiąga sufit (0.52 vs 0.50) — z zastrzeżeniem, że jego sufit jest miękki + (mała domena 348 melodii, sędzia najmniej pewny na 3/4). + +Ograniczenia: n=30/komórkę (±0.06) i sufit liczony na próbce, nie na całej puli — wersja +pod raport wymaga pełnego przebiegu (n=100) i referencji z całej puli domenowej. + +### Pliki + +| plik | rola | +|---|---| +| `src/core/abc_corpus.py` | wspólne czyszczenie ABC (reuse w prepare_data i judge) | +| `src/core/judge.py` | `JudgeGPT`: trunk GPT + mean-pool + 3 głowy | +| `src/tools/train_judge.py` | trening sędziego z tunes.csv | +| `src/tools/judge_tunes.py` | ocena plików .abc | +| `src/tools/benchmark_judge.py` | benchmark domenowy ekspertów (score/pokrycie/sufit) | +| `src/tools/benchmark_ood.py` | macierz OOD (ekspert × domena celu) | +| `run_benchmarks.sh` | sweep wszystkich checkpointów sędzią judge_v2 | +| `data/models/domains.json` | domena każdego checkpointu (type/meter) | +| `data/models/judge_v2.pt` | kanoniczny sędzia (5000 kroków) | +| `data/benchmarks/*.log` | logi przebiegów benchmarków | diff --git a/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Posttraining-ReverseKL-Eksperyment.md b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Posttraining-ReverseKL-Eksperyment.md new file mode 100644 index 0000000..4dffbd6 --- /dev/null +++ b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Posttraining-ReverseKL-Eksperyment.md @@ -0,0 +1,260 @@ +--- +type: nauka-log +title: "Posttraining: reverse KL z rotującym nauczycielem " +status: ZREALIZOWANY +data: 2026-09-05 +created_at: 2026-09-05 +author: Adam Skrodzki +tags: [nauka, llm, muzyka, capstone, gpt, ngram, realizacja] +repo_github: "https://github.com/adamskrodzki/micro-models" +--- + +# Posttraining: reverse KL z rotującym nauczycielem (eksperyment E-RKL) + +Status: ZREALIZOWANY (iteracja 2: continuation+scratch, α=0.1, reverse). Wyniki niżej. + +## Motywacja + +Macierz transferu OOD (`benchmark_ood.py`, log `ood_data_20260904_180711`) pokazuje: + +- **scratch**: eksperci są trenowani wyłącznie na jednym stylu, poza domeną wynik 0.07–0.15 + przy suficie 0.82/0.50, home-bias 0.89–0.95. Model ignoruje nagłówki + `M:`/`K:` i pisze w swoim metrum. +- **continuation**: zakotwiczenie promptem realnym ciałem melodii podwaja/ potraja wynik + poza domeną (reel 0.14→0.34, waltz 0.07→0.24), ale hb (home bias) spada tylko do ~0.4–0.55 — styl + domowy dalej przecieka. +- **Skala nie pomaga**: większe/dłużej trenowane ckpty (e128, e192) mają najlepszy wynik + w domenie i najwyższą sztywność (hb ~0.9–0.97 wszędzie). + +Wniosek: eksperci uczą się stylu domowego, ale nie uczą się *słuchać promptu*. Model +zachowuje się źle we własnym rozkładzie generacji — a dokładnie tam działa trening +on-policy. + +## Hipoteza + +On-policy reverse KL (student‖nauczyciel) z nauczycielem kompetentnym w domenie promptu +nauczy jeden checkpoint przełączać styl zależnie od nagłówków — obniży home-bias poza +domeną bez utraty jakości w domenie. W trakcie pracy pytanie zostało zaostrzone do: +**ile da się nauczyć „z drugiej ręki" — przez same prawdopodobieństwa nauczycieli, bez +widzenia surowych danych w fazie posttrainingu?** Metryka: +luka między studentem a universalistą (który te dane widział) w każdej komórce macierzy. + +## Metoda (stan po realizacji) + +- **Student**: `universalist_ckpt.pt` — GPT trenowany od zera na całym korpusie mieszanym + (`prepare_data.py --mixed`: 47k melodii, 10.5M znaków, 54-znakowy słownik, block 128, + 8000 iteracji ≈ 3 epoki). Początkowo tylko baseline „pierwszej ręki"; w toku pracy + stał się naturalnym punktem startowym posttrainingu (zastąpił pomysł startu z + eksperta — patrz Przebieg, krok 0). +- **Nauczyciele**: trzej specjaliści przetrenowani na WSPÓLNYM słowniku + (`VOCAB_FROM=universalist_ckpt.pt`): `jig_sh` (jig 6/8), `reel_sh` (reel 4/4), + `waltz_sh` (waltz 3/4), każdy 2000 iteracji na własnej domenie. +- **Zadania**: continuation (nagłówki + prefiks ciała o STAŁEJ długości 96 znaków — równe + długości promptów umożliwiają batchową generację) i scratch (same nagłówki). + Rotacja per batch: domena (jig→reel→waltz) × zadanie — 6 kombinacji w cyklu. +- **Pule promptów**: TRENINGOWE (split treningowy; val zostaje czysta dla ewaluacji), + filtr jak w benchmarku OOD: type substring + meter równość + ciało ≥ 160 znaków. +- **Rollout i loss**: student generuje 420 znaków (temp 0.85, top-k 18 — jak + w benchmarku); per-token KL liczony TYLKO na wygenerowanych pozycjach, w oknach 128 + ze stride 96, gdzie pierwsze 32 znaki okna to tylko KONTEKST (nauczyciel i student + przewidują z prawdziwego prefixu, nie z amputowanego środka melodii). +- **α-mixing**: `p_mix = (1−α)·p_teacher + α/|V|`, gdzie |V| = rozmiar wspólnego + słownika (54 znaki), α = 0.1. Czyli: rozkład nauczyciela zmiksowany w 10% z rozkładem + JEDNOSTAJNYM (α/|V| = 0.1/54 ≈ 0.0019 masy na każdy znak). Dwa efekty: + (a) **podłoga prawdopodobieństwa** — żaden znak nie ma w p_mix mniej niż α/|V|, więc + log p_mix jest ograniczony od dołu (≈ −6.3 nats); gradient KL nigdy nie eksploduje + na znakach, których nauczyciel nie zna (losowe logity, patrz Przebieg krok 1) ani tam, + gdzie jest pewny, a student generuje nie-nauczycielskie treści; + (b) **lekkie ściągnięcie studenta ku uniform** — kosztem: zbyt duże α spłaszcza + studenta (dlatego entropia jest metryką kontrolną, a sygnałem alarmowym jej wzrost + przy spadającym KL). Zastępuje ε-smoothing z pierwotnego planu — przy wspólnym + słowniku to to samo działanie zapisane inaczej. +- **Optymalizacja**: AdamW lr 1e-4 (fine-tune), warmup 100 + cosine, clip 1.0, bf16; + 4000 iteracji, batch 16. Ewaluacja co 200 iteracji: KL / tlp / entropia per + domena×zadanie na utwalonej puli val; zapisywane checkpointy best (min val KL) i last. + +## Ramy porównawcze + +| Arm | Trening | Status | +|---|---|---| +| baseline | universalista: SFT od zera na całym korpusie | zrobiony — punkt odniesienia „pierwszej ręki" i punkt startowy studenta | +| B | on-policy reverse KL, rotacja domena×zadanie | ZREALIZOWANY — wyniki niżej | +| A | mixed SFT fine-tune, off-policy CE | ODRZUCONY jako kontrola: widzi surowe dane domen, więc nie odpowiada na pytanie second-hand; jego rolę pełni baseline universalisty | +| A′ | off-policy distillation: KL α-miksowanego nauczyciela na PRAWDZIWYCH sekwencjach (teacher forcing) | zaproponowany — izoluje on-policy przy identycznym zestawie informacji | +| C | on-policy forward KL (`--direction forward`), ta sama rotacja | otwarty — izoluje kierunek KL | + +## Metryki + +- **Macierz OOD** (`benchmark_ood.py`, ta sama konfiguracja co log bazowy): score, ref, + hb, pokrycie — scratch i continuation. +- **Benchmark domenowy** własnego checkpointu: zapominanie domeny domowej. +- **Trajektoria hb** w trakcie treningu — bezpośredni odczyt uczenia się słuchania promptu. +- **Entropia generacji per domena** — wykrywanie kolapsu różnorodności wewnątrz domeny + (rotacja chroni różnorodności między domenami, ale nie wewnątrz). +- **Sanity-check na start**: średni log-prob nauczyciela na rolloutach studenta dla każdej + pary (student × nauczyciel). Jeśli nauczyciel waltz daje katastrofalnie niski log-prob + na śmieciach jig-studenta, gradienty będą zdominowane przez pozycje „unikaj niskiego + prawdopodobieństwa" zamiast „bądź jak waltz" — przerwać i przemyśleć (np. maskowanie + pozycji, mocniejsze wygładzenie). + +## Kryteria sukcesu / dyskryminatory + +Po zmianie punktu startowego na universalistę kryteria przeformułowane — porównanie +z ekspertem-startem straciło sens, liczy się luka do universalisty (first-hand) i do +nauczycieli (specialist ceiling). RKL musiało: + +1. **Podnieść komórki poza domeną** — zwłaszcza waltz (najsłabsza: 0.20/0.29), +2. **Nie pogorszyć jig** (zapominanie domeny największej jakościowo), +3. **Obniżyć home bias poza domeną** (słuchanie nagłówków). + +Wszystkie trzy spełnione w iteracji 2 — patrz Wyniki i Wnioski. + +## Znane ryzyka + +- **Mode-seeking reverse KL** → wrażliwość studenta na argmax nauczyciela; monitorować + entropię per domena. +- **Nauczyciel nieinformatywny na rolloutach studenta** poza domeną — sanity-check wyżej. +- **Mała pula waltz** (348) → resampling, efektywna repetycja; ewentualna augmentacja. +- **Scratch może się nie poprawić** — trening tylko na continuation; hb w scratch może + zostać wysokie (brak zakotwiczenia). Ewaluować oba zadania, nie zakładać transferu. +- **Kolaps dywersyjności wewnątrz domeny** mimo rotacji. + +--- + +## Przebieg (chronologicznie) + +Implementacja: `src/train/train_rkl.py`. Kroki i decyzje w kolejności, w jakiej +zapadały: + +0. **Baseline: universalista.** Trening mixed SFT na pełnym korpusie (10.5M znaków, + 8000 iteracji) i benchmark OOD: przełącza style (reel scratch 0.48 vs 0.07–0.15 + ekspertów), ale w każdej domenie płytko (jig 0.52/0.62, waltz 0.20/0.29). Ustalony + jako punkt odniesienia first-hand — i, w trakcie dyskusji nad punktem startowym + posttrainingu (kandydaci pierwotni: `jig_e128_s3` vs `jig_e32_s3`), zarzucony + na rzecz universalisty jako studenta: start z eksperta dawałby rigidity do zdjęcia, + a start z universalisty testuje czysty przyrost „na drugą rękę". +1. **Problem słownika (decyzja architektoniczna).** Modele są char-level: słownik = + tokenizer = `sorted(set(korpusu))`, każdy checkpoint ma własny. Dwa poziomy problemu: + (a) ten sam znak ma inne ID w różnych ckpt — trywialne, remap po znaku; (b) znak + nieobecny w korpusie nauczyciela NIE ma wytrenowanego embeddingu ani logitu + (weight tying: nieużywane wiersze nie dostają gradientu) — nauczyciel nie potrafi + go reprezentować, a jego logit to szum inicjalizacji. Rozważane konwencje: + renormalizacja do supportu nauczyciela / ε-floor / pomijanie pozycji. Decyzja: rozwiązać + u źródła — JEDEN wspólny tokenizer z pełnego zbioru (54 znaki), przetrenowanie + nauczycieli z `VOCAB_FROM=universalist_ckpt.pt` (`*_sh_ckpt.pt`). To usuwa problem + reprezentowalności, ale nie kalibracji: znak wspólnego słownika niewystępujący + w korpusie nauczyciela nadal ma losowy logit. Stąd α-mixing zamiast ε-floor — + przy wspólnym słowniku to to samo działanie zapisane inaczej (definicja i mechanika + w Metodzie), a dodatkowo ogranicza gradienty tam, gdzie nauczyciel jest pewny, + a student (jeszcze) generuje nie-nauczycielskie treści. +2. **Nauczyciele `_sh`: sanity benchmark.** Diagonale nie gorsze niż u starych ekspertów + (jig 0.64, reel 0.57, waltz 0.53 — reel i waltz lepiej), pokrycie 100% (szumne logity + niewytrenowanych wierszy nie psują generacji), rigidity zachowana (hb 0.9+ poza + domeną). Zweryfikowani jako supervisorzy. +3. **Iteracja 1: RKL continuation-only.** Pierwsza wersja ciąła sekwencje na okna 128 + bez zachodzenia — nauczyciel oceniał środki melodii bez prawdziwego poprzedzenia, + jego rozkłady były rozmyte i student SPŁASZCZAŁ SIĘ do uniform (entropia 2.2→3.0 + nats/znak przy spadającym KL: mechanizm „tania zgodność z rozmytym celem"). Poprawka: + okna ze stride 96, pierwsze 32 znaki okna to tylko kontekst — entropia wróciła do + ~1.0, tlp (log-prob nauczyciela na rolloutach studenta) zaczęło rosnąć. Zbiegł do + val KL 0.155. Benchmark: waltz 0.20→0.37/0.43, reel cont 0.52→0.56, jig bez zmian — + ALE reel scratch 0.48→0.26: tryb nigdy nieobecny w treningu zregresował. +4. **Iteracja 2: dodanie zadania scratch.** Rotacja domena×zadanie (6 kombinacji + w cyklu batchy). Val KL 0.117, entropia stabilna ~1.3–1.4, tlp najlepsze z przebiegu. + Benchmark: reel scratch naprawiony (0.26→0.53), wszystkie komórki ≥ universalisty — + Pareto-dominacja (pełna tabela niżej). +5. **Uwaga metodologiczna: best-by-KL słabo skorelowany z jakością benchmarkową** — + `last` był lepszy na komórkach istotnych (scratch) mimo wyższego val KL; oba + checkpointy zapisywane od tej pory. +6. **Ablacja: start z eksperta zamiast z universalisty.** Pytanie: ile wyniku bierze się + z faktu, że universalista widział wszystkie domeny PRZED RKL? Start: `jig_sh` (specjalista, + nigdy nie widział reel/waltz na żadnym etapie) + ci sami trzej nauczyciele, ten sam + budżet (`student_jigstart_*`). Przebieg zdrowy (KL 0.146, entropia stabilna); reel + domykał się najwolniej (KL 0.32 vs 0.065 waltz). Benchmark patrz tabela ablacji + w Wynikach. + +## Wyniki (judge_v2, seed 42; ref: jig 0.799 / reel 0.766 / waltz 0.544) + +Macierz OOD, komórka = `scratch` / `continuation`: + +| model | jig:6/8 | reel:4/4 | waltz:3/4 | +|---|---|---|---| +| nauczyciel jig_sh (diag) | 0.64 / 0.72 | 0.14 / 0.31 | 0.16 / 0.23 | +| nauczyciel reel_sh (diag) | 0.25 / 0.43 | 0.57 / 0.59 | 0.16 / 0.21 | +| nauczyciel waltz_sh (diag) | 0.14 / 0.38 | 0.12 / 0.33 | 0.53 / 0.45 | +| universalist (baseline) | 0.52 / 0.62 | 0.48 / 0.52 | 0.20 / 0.29 | +| RKL continuation-only | 0.50 / 0.63 | 0.26 / 0.56 | 0.37 / 0.43 | +| **RKL cont+scratch (last)** | **0.62 / 0.66** | **0.53 / 0.57** | **0.54 / 0.43** | + +(off-diagonal nauczycieli z `ood_teachers_*.log`; checkpoint continuation-only usunięty +przy sprzątaniu — liczby z logu `ood_student_rkl_*.log`) + +### Ablacja: universalist-start vs specialist-start + +Ten sam trening (B), inny punkt startowy — student nigdy nie widzący reel/waltz +na żadnym etapie (`student_jigstart_last`, log `ood_jigstart_*.log`): + +| komórka (scr / cont) | universalist | RKL uni-start | **RKL jig-start** | jig_sh (init/nauczyciel) | +|---|---|---|---|---| +| jig | 0.52 / 0.62 | 0.62 / 0.66 | 0.60 / 0.63 | 0.62 / 0.70 | +| reel | 0.48 / 0.52 | 0.53 / 0.57 | 0.39 / 0.53 | 0.15 / 0.31 | +| waltz | 0.20 / 0.29 | 0.54 / 0.43 | 0.55 / 0.44 | 0.14 / 0.22 | + +Wniosek z ablacji: **pre-RKL ekspozycja na dane wszystkich domen wnosi niemal nic** — +waltz i reel-continuation identyczne między startami; jedyna różnica to reel scratch +(0.39 vs 0.53), czyli tryb, w którym model musi wygenerować cały styl z samego priora, +bez melodicznego zakotwiczenia — dokładnie tam pierwszy kontakt z prawdziwymi danymi +pomaga, a nadzór prawdopodobieństwami najmniej. Dodatkowo: rigidity specjalisty jest +w pełni zdejmowalna (hb 0.94 → 0.04–0.25), a cena za domyk w dwóch obcych domenach to +lekki spadek w domenie własnej (jig 0.63 vs 0.70 nauczyciela w continuation). + +**Wnioski:** + +1. **Second-hand działa — i to prawie w całości**: student po RKL zbliża się do każdego + specjalisty na jego własnej przekątnej (0.02–0.04), trenowany w fazie RKL tylko na + prawdopodobieństwach nauczycieli; ablacja (start z eksperta bez kontaktu z danymi + reel/waltz w ogóle) pokazuje, że pre-RKL ekspozycja universalisty wnosi wyłącznie + ~0.14 na komórce reel scratch. +2. **Waltz: 0.20→0.54 scratch** — sufit sędziego (0.544), poziom własnego nauczyciela + (0.53). Najmniejsza domena (pula 348) domknięta. +3. **Pareto-dominacja nad universalistą** w każdej komórce — brak kosztu zapominania; + off-diagonal hb spadło do 0.06–0.15 (student słucha nagłówków we wszystkich kierunkach). +4. **Scratch musi być w treningu**, żeby nie zregresować (Przebieg, krok 4). +5. **Rigidity ekspertów nie przeniosła się na studenta**: nauczyciele poza domeną + generują w swoim metrum (hb 0.9+), rotacja trzech nauczycieli nauczyła studenta + przełączania zależnie od nagłówków zamiast jednego sztywnego stylu. +6. Sanitarny wskaźnik w trakcie treningu: tlp (log-prob nauczyciela na rolloutach + studenta) rosnący + stabilna entropia = zdrowy przebieg; rosnąca entropia przy + spadającym KL = spłaszczanie (sygnał do przerwania / zmiany α). + +**Otwarte:** forward KL (`--direction forward`) — izolacja kierunku KL; off-policy +distillation (nauczyciel na prawdziwych sekwencjach) — izolacja on-policy; sensitivity +α i top-k; odtworzenie checkpointu continuation-only dla pełnej tabeli abacyjnej. + +**Zastrzeżenie (leak generatory↔val sędziego):** sędzia nie wycieka do treningu w żadnej +formie (w `train_rkl.py` z `train_judge` importowane są tylko narzędzia danych; sygnał +treningowy to wyłącznie rozkłady nauczycieli; pule promptów RKL z splitu TRENINGOWEGO). +Ale generatory (universalista, nauczyciele `_sh`) trenowały na pełnym `tunes.csv`, +w tym na melodiach z walidacyjnego splitu sędziego — benchmark może więc mierzyć częściowo +zapamiętywanie (najbardziej continuation: prompt = ćwiartka ciała melodii walidacyjnej; +najmniej scratch). Wyciek wspólny dla wszystkich porównywanych modeli, więc wnioski +relatywne (RKL vs universalista, uni-start vs jig-start) pozostają w mocy; wartości +absolutne traktować jako optymistyczne. Clean-room wariant i szczegóły: [[Judge-Sedzia-Generacji]], +sekcja „Zastrzeżenie: split chroni sędziego, ale nie generatory". + +### Konsekwencje: co zostaje, a co nie + +Wyciek działa jak stała dopłatka do wyników każdego generatora, więc: + +- **Porównania (różnice) zostają**: RKL vs universalista, uni-start vs jig-start, + nauczyciel vs student — wszystkie modele dziedziczą ten sam nadmiar informacji, + więc różnice między nimi mierzą realny efekt interwencji, nie wyciek. +- **Ratia score/ref zostają**: sufit (`ref`) liczony na prawdziwych melodiach jest + niewrażliwy na zapamiętywanie przez generatory; normalizacja sufitem częściowo + usuwa wspólną dopłatkę. +- **Wartości absolutne nie zostają**: score 0.54 na waltz to górne oszacowanie — + część tej liczby to odtworzenie zapamiętanych ciał, nie opanowanie stylu. +- **Granica zaufania**: dopłatka nie musi być identyczna między modelami (pojemność + i liczba epok na korpusie z wyciekiem różnią się między checkpointami) — wnioski + oparte na dużych lukach (≥0.1, np. waltz 0.20→0.54) są bezpieczne; porównania na + granicy szumu (±0.03–0.05) przy wycieku tracą jeszcze trochę wiarygodności. diff --git a/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Reprodukcja-E-RKL-Krok-Po-Kroku.md b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Reprodukcja-E-RKL-Krok-Po-Kroku.md new file mode 100644 index 0000000..295aa6b --- /dev/null +++ b/music-experts/docs/Badania/2026-09-05_posttraining-reverse-kl/Reprodukcja-E-RKL-Krok-Po-Kroku.md @@ -0,0 +1,193 @@ +--- +type: reprodukcja +title: "E-RKL — reprodukcja krok po kroku" +status: aktywne +data: 2026-09-05 +created_at: 2026-09-05 +author: Adam Skrodzki +tags: [nauka, llm, muzyka, gpt, realizacja, reprodukcja] +repo_github: "https://github.com/adamskrodzki/micro-models" +--- + +# E-RKL — reprodukcja krok po kroku + +Jak odtworzyć eksperyment „Posttraining: reverse KL z rotującym nauczycielem" od surowych +danych po finalne benchmarki. Kontekst i wyniki: [[Posttraining-ReverseKL-Eksperyment]]. +Wszystkie komendy uruchamiamy z katalogu `music-experts/`. Zalecany CUDA (treningi i +benchmarki na CPU działają, ale wolno). + +## Wymagania wstępne + +- `data/tunes.csv` — dane źródłowe, format jak w repo - oryginalne kolumny `tune_id,type,meter,mode,abc`). +- `data/models/judge_v2.pt` — sędzia (klasyfikator meter/mode/type). Jeśli brak: + `python src/tools/train_judge.py --csv data/tunes.csv --out data/models/judge_v2.pt` + (seed splitu 42 — benchmarki zakładają ten sam `--split-seed`). +- `pip install -r requirements.txt`; generator liczb losowych i split są deterministyczne + (seed 42 wszędzie), więc wyniki powinny być odtwarzalne do szumu generacji. + +## Krok 1 — korpusy ABC + +Trzy korpusy domenowe + jeden mieszany (wszystko z `tunes.csv`, normalizacja przez +`clean_abc`/`norm_key`): + +```bash +python src/data/prepare_data.py jig 6/8 data/jigs.abc +python src/data/prepare_data.py reel 4/4 data/reels.abc +python src/data/prepare_data.py waltz 3/4 data/waltzes.abc +python src/data/prepare_data.py --mixed data/mixed.abc +``` + +Oczekiwane rzędy wielkości (wersja tunes.csv z 2026-09): jigi ~2.5M znaków, mixed +~10.5M znaków, 54 znaki słownika. `--mixed` nie filtruje typu/metrum — nagłówek `M:` +bierze z wiersza. + +## Krok 2 — universalista (wspólny słownik + baseline first-hand) + +```bash +python src/train/train_gpt.py data/mixed.abc \ + data/models/universalist_ckpt.pt data/models/universalist_loss_log.csv \ + "" 20260620 --max-iters 8000 +``` + +Parametry domyślne: block 128, batch 32, lr 3e-4, N_LAYER/N_EMBD z ENV (tu: domyślne 4/128). +Kluczowe: to trenowanie BUDUJE słownik 54 znaków — wszyscy nauczyciele w kroku 3 +startują z `VOCAB_FROM` na tym ckpcie. + +## Krok 3 — nauczyciele na wspólnym słowniku + +```bash +python src/train/train_gpt.py data/jigs.abc data/models/jig_sh_ckpt.pt data/models/jig_sh_loss.csv data/models/universalist_ckpt.pt +python src/train/train_gpt.py data/reels.abc data/models/reel_sh_ckpt.pt data/models/reel_sh_loss.csv data/models/universalist_ckpt.pt +python src/train/train_gpt.py data/waltzes.abc data/models/waltz_sh_ckpt.pt data/models/waltz_sh_loss.csv data/models/universalist_ckpt.pt +``` + +Czwarty argument pozycyjny to `VOCAB_FROM` — ładuje `stoi`/`itos` z universalisty +(wspólny tokenizer; patrz Przebieg krok 1 w [[Posttraining-ReverseKL-Eksperyment]]). +Przy 2000 iteracji każda komenda wypisuje val loss co 200 i zapisze best ckpt. + +## Krok 4 — rejestr domen + +Dopisz do `data/models/domains.json` (pole `type` decyduje o puli benchmarku i gwiazdce +diagonalnej; dla modeli bez „domeny domowej" wpis jest arbitralny): + +```json +"universalist_ckpt.pt": {"type": "jig", "meter": "6/8"}, +"jig_sh_ckpt.pt": {"type": "jig", "meter": "6/8"}, +"reel_sh_ckpt.pt": {"type": "reel", "meter": "4/4"}, +"waltz_sh_ckpt.pt": {"type": "waltz", "meter": "3/4"} +``` + +JSON musi być poprawny (`python -c "import json; json.load(open('data/models/domains.json'))"`) +— benchmarki bez zapasu kończą się wyjątkiem parsowania. + +## Krok 5 — sanity benchmark nauczycieli i baseline universalisty + +```bash +python src/tools/benchmark_ood.py \ + --models data/models/jig_sh_ckpt.pt data/models/reel_sh_ckpt.pt data/models/waltz_sh_ckpt.pt \ + | tee data/benchmarks/ood_teachers_$(date +%Y%m%d_%H%M%S).log + +python src/tools/benchmark_ood.py --models data/models/universalist_ckpt.pt \ + | tee data/benchmarks/ood_universalist_$(date +%Y%m%d_%H%M%S).log +``` + +Oczekiwania: nauczyciele — wysokie diagonale (jig ~0.64, reel ~0.57, waltz ~0.53), +pokrycie 100% w scratch, home bias ~0.9 poza domeną; universalista — przełącza style (reel +scratch ~0.48), ale płytko (waltz ~0.20). Te logi są punktem odniesienia dla kroku 8. + +## Krok 6 — trening RKL (pomysł główny) + +Najpierw smoke test (mechanika, kształty, zapis — NIE nadpisuj `--last` z właściwego +runu!): + +```bash +python src/train/train_rkl.py \ + --teachers jig=data/models/jig_sh_ckpt.pt reel=data/models/reel_sh_ckpt.pt waltz=data/models/waltz_sh_ckpt.pt \ + --max-iters 20 --batch-size 4 --eval-interval 10 --eval-samples 4 \ + --out data/models/_smoke.pt --losslog data/models/_smoke_loss.csv +rm data/models/_smoke.pt data/models/_smoke_loss.csv +``` + +Pełny run (α=0.1, reverse, obie rotacje: domena×zadanie): + +```bash +python src/train/train_rkl.py \ + --teachers jig=data/models/jig_sh_ckpt.pt reel=data/models/reel_sh_ckpt.pt waltz=data/models/waltz_sh_ckpt.pt \ + --max-iters 4000 --batch-size 16 \ + --out data/models/student_rkl_sc_ckpt.pt --last data/models/student_rkl_sc_last_ckpt.pt \ + --losslog data/models/student_rkl_sc_loss.csv +``` + +Jak czytać przebieg (co 200 iteracji, per domena×zadanie): +- `kl` — spada (cel: < ~0.2 na startcie z universalisty), +- `tlp` — log-prob nauczyciela na rolloutach studenta; ma ROSNĄĆ (nauczyciel coraz lepiej + rozumie studenta), +- `ent` — entropia generacji; ma być STABILNA (~1.0–1.8 akceptowalne, idealnie < 1.5). Rosnąca entropia przy + spadającym KL = spłaszczanie do uniform → przerwać, zmniejszyć `--alpha` (np. 0.03). + +## Krok 7 — benchmarki finalne + +Dopisz checkpointy studenta do `domains.json` (dowolna istniejąca domena — pole służy +tylko jako „home" do hb i gwiazdki): + +```json +"student_rkl_sc_ckpt.pt": {"type": "jig", "meter": "6/8"}, +"student_rkl_sc_last_ckpt.pt": {"type": "jig", "meter": "6/8"} +``` + +```bash +python src/tools/benchmark_judge.py --model data/models/student_rkl_sc_ckpt.pt --samples 100 \ + | tee data/benchmarks/bench_student_rkl_$(date +%Y%m%d_%H%M%S).log + +python src/tools/benchmark_ood.py --models data/models/student_rkl_sc_ckpt.pt data/models/student_rkl_sc_last_ckpt.pt \ + | tee data/benchmarks/ood_student_sc_$(date +%Y%m%d_%H%M%S).log +``` + +Referencje do porównania (seed 42, judge_v2): uniwersalista 0.52/0.62 · 0.48/0.52 · +0.20/0.29; nauczyciele na przekątnych 0.64 · 0.57 · 0.53 (scratch). Oczekiwany wynik: +student Pareto-dominuje universalistę, waltz scratch na suficie sędziego (~0.54). + +## Krok 8 (opcjonalny) — ablacja startu z eksperta + +Ten sam trening, ale student startuje z nauczyciela jig (nigdy nie widział reel/waltz +na żadnym etapie) — izoluje wkład pre-RKL ekspozycji universalisty: + +```bash +python src/train/train_rkl.py \ + --student data/models/jig_sh_ckpt.pt \ + --teachers jig=data/models/jig_sh_ckpt.pt reel=data/models/reel_sh_ckpt.pt waltz=data/models/waltz_sh_ckpt.pt \ + --max-iters 4000 --batch-size 16 \ + --out data/models/student_jigstart_ckpt.pt --last data/models/student_jigstart_last_ckpt.pt \ + --losslog data/models/student_jigstart_loss.csv +``` + +(Przy `--student` innym niż universalista asercja słownika musi przejść — dlatego +nauczyciele muszą być z kroku 3, nie stare eksperckie ckpty.) Benchmark jak w kroku 7. + +## Artefakty + +| ścieżka | co | +|---|---| +| `data/jigs.abc`, `data/reels.abc`, `data/waltzes.abc`, `data/mixed.abc` | korpusy ABC | +| `data/models/universalist_ckpt.pt` | student startowy / wspólny słownik (54 znaki) | +| `data/models/{jig,reel,waltz}_sh_ckpt.pt` | nauczyciele na wspólnym słowniku | +| `data/models/student_rkl_sc_{,_last_}ckpt.pt` | student po RKL (best-by-KL / last) | +| `data/models/student_jigstart_{,_last_}ckpt.pt` | ablation: start z jig_sh | +| `data/models/domains.json` | rejestr domen benchmarku | +| `data/models/*_loss.csv`, `data/benchmarks/*.log` | krzywe treningu i logi benchmarków | +| `src/data/prepare_data.py`, `src/train/train_gpt.py`, `src/train/train_rkl.py` | pipeline | +| `src/tools/{train_judge,benchmark_judge,benchmark_ood}.py` | sędzia i benchmarki | + +## Pułapki + +1. **Słownik**: `train_rkl.py` wymaga identycznego zestawu znaków studenta i nauczycieli + (twarda asercja). Stare eksperckie ckpty mają własne, mniejsze słowniki — nie działają + jako nauczyciele ani (bez przetrenowania) jako student. +2. **`--last`/`--out`**: zawsze podawaj jawne ścieżki; domyślne nadpisują się między + runami (zdarzyło się: smoke test nadpisał `student_rkl_last.pt`). +3. **Pokrycie w benchmarkach**: pominięte generacje liczą się jako 0; porównuj modele + tylko przy zbliżonym pokryciu (spadki pokrycia na continuation przy waltz — 83% — + wynikają ze znaków w surowych wierszach CSV spoza słownika, nie z jakości modelu). +4. **Seed**: benchmarki są deterministyczne przy tym samym seed (identyczne melodie + w kolumnach); generacja przy ewaluacji i tak wnosi szum ±0.03–0.05 — nie czytaj + pojedynczych eval-i zbyt dosłownie. diff --git a/music-experts/docs/Badania/Badania-INDEX.md b/music-experts/docs/Badania/Badania-INDEX.md index 7ee4dbf..6694a7e 100644 --- a/music-experts/docs/Badania/Badania-INDEX.md +++ b/music-experts/docs/Badania/Badania-INDEX.md @@ -19,6 +19,7 @@ repo: "https://github.com/slayerlabs/micro-models" | 2026-06-20 | Kompozycja małych modeli (kontrakt · mapper · wariancja) | warsztat dialektyczny | [[Kompozycja-INDEX]] | | 2026-06-20 | N-gram → mini-transformer (most) | koncepcja / roadmap | [[Research-NGram-vs-MiniTransformer]] | | 2026-06-21 | Granie razem — sprzężone oscylatory + polifonia (oś pionowa) | koncepcja | [[Granie-Razem-Polifonia]] | +| 2026-09-05 | Posttraining: reverse KL z rotującym nauczycielem (E-RKL) | zrealizowany | [[Posttraining-ReverseKL-Eksperyment]] · [[Reprodukcja-E-RKL-Krok-Po-Kroku]] | ## Kotwica strategiczna - [[Cele-Globalne-i-Kotwica]] — hierarchia: muzyka = sandbox (Cel #1 know-how) · IFC = produkt (Cel #2). Co przenosi się uczciwie, kolejność budowy IFC, walidator = nagroda RLVR. diff --git a/music-experts/renderer.py b/music-experts/renderer.py new file mode 100644 index 0000000..2d7b0a9 --- /dev/null +++ b/music-experts/renderer.py @@ -0,0 +1,41 @@ +# Autor: Adam Skrodzki +"""Render MIDI -> WAV (sine-wave synthesis, no external binaries needed). +Usage: python renderer.py input.mid [output.wav] +""" +import sys +import numpy as np +import pretty_midi +import wave + +def main(): + if len(sys.argv) < 2: + sys.exit("usage: python renderer.py input.mid [output.wav]") + mid_path = sys.argv[1] + wav_path = sys.argv[2] if len(sys.argv) > 2 else mid_path.rsplit(".", 1)[0] + ".wav" + + pm = pretty_midi.PrettyMIDI(mid_path) + instruments = pm.instruments + note_counts = [len(i.notes) for i in instruments] + duration = pm.get_end_time() + print(f"in: {mid_path}") + print(f" instruments: {len(instruments)} | notes per instrument: {note_counts} | duration: {duration:.1f}s") + + if sum(note_counts) == 0: + sys.exit("ERROR: MIDI contains no notes — nothing to render. Check the generation/conversion step.") + + audio = pm.synthesize(fs=44100) + peak = float(np.abs(audio).max()) + print(f" peak amplitude before normalize: {peak:.4f}") + if peak == 0: + sys.exit("ERROR: synthesis produced silence despite notes present.") + + audio = (audio / peak * 32767).astype(np.int16) + with wave.open(wav_path, "w") as w: + w.setnchannels(1) + w.setsampwidth(2) + w.setframerate(44100) + w.writeframes(audio.tobytes()) + print(f"out: {wav_path} | {len(audio) / 44100:.1f}s @ 44.1kHz mono int16") + +if __name__ == "__main__": + main() diff --git a/music-experts/run_benchmarks.sh b/music-experts/run_benchmarks.sh new file mode 100755 index 0000000..b48a76c --- /dev/null +++ b/music-experts/run_benchmarks.sh @@ -0,0 +1,26 @@ +#!/usr/bin/env bash +# Autor: Adam Skrodzki +# Benchmark wszystkich ekspertów GPT (data/models/*_ckpt.pt) sędzią judge_v2. +# Log każdej sesji do data/benchmarks/bench_.log — wyniki interpretujemy później. +# Użycie: ./run_benchmarks.sh [judge] [samples] +# ./run_benchmarks.sh # judge_v2, 100 próbek +# ./run_benchmarks.sh data/models/judge_v2.pt 50 # inny sędzia / mniej próbek +set -u +cd "$(dirname "$0")" + +JUDGE="${1:-data/models/judge_v2.pt}" +SAMPLES="${2:-100}" +OUT="data/benchmarks/bench_$(date +%Y%m%d_%H%M%S).log" +mkdir -p "$(dirname "$OUT")" + +echo "sędzia: $JUDGE | próbek/model: $SAMPLES | log: $OUT" | tee "$OUT" +for ckpt in data/models/*_ckpt.pt; do + [ -e "$ckpt" ] || continue # glob bez trafień + echo "================================================================" | tee -a "$OUT" + echo ">>> $ckpt" | tee -a "$OUT" + python src/tools/benchmark_judge.py --model "$ckpt" --judge "$JUDGE" \ + --samples "$SAMPLES" 2>&1 | tee -a "$OUT" +done + +echo "================================================================" | tee -a "$OUT" +echo "gotowe: $(ls data/models/*_ckpt.pt | wc -l) checkpointów -> $OUT" | tee -a "$OUT" diff --git a/music-experts/src/core/abc_corpus.py b/music-experts/src/core/abc_corpus.py new file mode 100644 index 0000000..0beebf4 --- /dev/null +++ b/music-experts/src/core/abc_corpus.py @@ -0,0 +1,42 @@ +"""Wspólne narzędzia czyszczenia ABC (używane przez prepare_data.py i judge). +Przeniesione 1:1 z prepare_data.py + parser ciała melodii dla judge'a. +""" +import re + +MODE_TABLE = { + "major": "", "ionian": "", "minor": "min", "aeolian": "min", + "dorian": "dor", "mixolydian": "mix", "phrygian": "phr", + "lydian": "lyd", "locrian": "loc", "": "", +} + +ALLOWED = set("ABCDEFGabcdefg0123456789|:[]()<>/'^_=.,~- zZxX") + +HEADER_PREFIXES = ("X:", "T:", "C:", "M:", "K:", "N:", "%") + + +def norm_key(mode: str) -> str: + m = re.match(r"^([A-Ga-g][#b]?)(.*)$", mode.strip()) + if not m: + return "C" + root, word = m.group(1), m.group(2).lower() + return root + MODE_TABLE.get(word, "") + + +def clean_abc(body: str) -> str: + body = body.replace("\r\n", "\n").replace("\r", "\n").strip() + body = re.sub(r'"[^"]*"', "", body) # usuń symbole akordów / adnotacje "..." + body = re.sub(r"[ \t]+", " ", body) # scal podwójne spacje po usunięciu + body = re.sub(r"\n+", "\n", body) + return body + + +def tune_body(abc: str) -> str: + """Ciało melodii bez linii nagłówkowych (X:/T:/C:/M:/K:/N:/%).""" + body = "\n".join(ln for ln in abc.split("\n") if not ln.startswith(HEADER_PREFIXES)) + return clean_abc(body) + + +def parse_tune_row(row: dict) -> tuple[str, str, str, str]: + """(body, meter, mode_label, type_label) z wiersza tunes.csv; mode normalizowane.""" + body = tune_body(row["abc"]) + return body, row["meter"].strip(), norm_key(row["mode"]), row["type"].strip().lower() diff --git a/music-experts/src/core/judge.py b/music-experts/src/core/judge.py new file mode 100644 index 0000000..8b11957 --- /dev/null +++ b/music-experts/src/core/judge.py @@ -0,0 +1,46 @@ +# Autor: Adam Skrodzki +"""Judge — klasyfikator melodii ABC na trunku GPT (verbatim reuse core/gpt.py). +Trunk = cały GPT (embedding -> bloki -> ln_f); nowa głowa = pooling po pozycjach +(mean albo attention-pool, pool="mean"|"attn") + liniowe głowy klasyfikacyjne +(meter/mode/type). +Uwaga przyczynowa: padding na KOŃCU sekwencji — tokeny prawdziwe nigdy go nie widzą, +maska potrzebna tylko do poprawnego mean-poolingu. +""" +import torch +import torch.nn as nn +from core.gpt import GPT, GPTConfig + +HEADS = ("meter", "mode", "type") + + +class JudgeGPT(nn.Module): + def __init__(self, cfg: GPTConfig, n_classes: dict, pool: str = "mean"): + super().__init__() + self.gpt = GPT(cfg) # trunk verbatim (LM-head nieużywany) + self.pool = pool + # ModuleList zamiast ModuleDict: klucz "type" koliduje z nn.Module.type() + self.heads = nn.ModuleList([nn.Linear(cfg.n_embd, n_classes[h]) for h in HEADS]) + if pool == "attn": + # uwaga pooling: wagi uczące się per głowa; przy zerowych wagach = mean-pool + self.score = nn.Linear(cfg.n_embd, len(HEADS)) + + def trunk(self, idx): + g = self.gpt + pos = torch.arange(idx.shape[1], device=idx.device) + x = g.drop(g.tok_emb(idx) + g.pos_emb(pos)) + for blk in g.blocks: + x = blk(x) + return g.ln_f(x) # (B, T, C) + + def forward(self, idx, mask): + """idx: (B,T) | mask: (B,T) — 1 = prawdziwy znak, 0 = padding (tylko na końcu). + Zwraca {head: logity (B, n_classes)}.""" + h = self.trunk(idx) + if self.pool == "attn": + s = self.score(h).masked_fill(mask.unsqueeze(-1) == 0, float("-inf")) # (B,T,H) + w = torch.softmax(s, dim=1) # wagi po T + pooled = torch.einsum("bth,btc->bhc", w, h) # (B,H,C) + return {head: self.heads[k](pooled[:, k]) for k, head in enumerate(HEADS)} + m = mask.unsqueeze(-1).to(h.dtype) # (B,T,1) + pooled = (h * m).sum(dim=1) / m.sum(dim=1).clamp(min=1) # mean-pool po prawdziwych znakach + return {head: self.heads[k](pooled) for k, head in enumerate(HEADS)} diff --git a/music-experts/src/data/prepare_data.py b/music-experts/src/data/prepare_data.py index b697cf8..d3fd152 100644 --- a/music-experts/src/data/prepare_data.py +++ b/music-experts/src/data/prepare_data.py @@ -1,65 +1,57 @@ """Przygotowanie korpusu ABC z thesession.org tunes.csv. Filtr: typ + metrum z argumentów (domyślnie jigi 6/8). Buduje grające bloki ABC + normalizuje tonację. -Użycie: python src/data/prepare_data.py [typ] [metrum] [wyjście] +--mixed: bez filtra typu/metrum (uniwersalista, nagłówek M: z wiersza). +Użycie: python src/data/prepare_data.py [typ] [metrum] [wyjście] [--mixed] np. python src/data/prepare_data.py waltz 3/4 data/corpus/waltz.abc + python src/data/prepare_data.py --mixed data/mixed.abc """ -import csv, re, sys, os -csv.field_size_limit(10**7) - -TYPE_KW = sys.argv[1] if len(sys.argv) > 1 else "jig" -METER = sys.argv[2] if len(sys.argv) > 2 else "6/8" -OUT = sys.argv[3] if len(sys.argv) > 3 else "data/jigs.abc" - -MODE_TABLE = { - "major": "", "ionian": "", "minor": "min", "aeolian": "min", - "dorian": "dor", "mixolydian": "mix", "phrygian": "phr", - "lydian": "lyd", "locrian": "loc", "": "", -} - -def norm_key(mode: str) -> str: - m = re.match(r"^([A-Ga-g][#b]?)(.*)$", mode.strip()) - if not m: - return "C" - root, word = m.group(1), m.group(2).lower() - return root + MODE_TABLE.get(word, "") +import csv, sys, os, argparse +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from core.abc_corpus import ALLOWED, clean_abc, norm_key -ALLOWED = set("ABCDEFGabcdefg0123456789|:[]()<>/'^_=.,~- zZxX") +csv.field_size_limit(10**7) -def clean_abc(body: str) -> str: - body = body.replace("\r\n", "\n").replace("\r", "\n").strip() - body = re.sub(r'"[^"]*"', "", body) # usuń symbole akordów / adnotacje "..." - body = re.sub(r"[ \t]+", " ", body) # scal podwójne spacje po usunięciu - body = re.sub(r"\n+", "\n", body) - return body +ap = argparse.ArgumentParser() +ap.add_argument("type_kw", nargs="?", default="jig") +ap.add_argument("meter", nargs="?", default="6/8") +ap.add_argument("out", nargs="?", default=None) +ap.add_argument("--mixed", action="store_true", + help="cały korpus bez filtra typu/metrum (M: z wiersza)") +a = ap.parse_args() +if a.out is None: + a.out = "data/mixed.abc" if a.mixed else "data/jigs.abc" def main(): rows_out, n_total, n_kept = [], 0, 0 with open("data/tunes.csv", encoding="utf-8", newline="") as f: for row in csv.DictReader(f): n_total += 1 - if TYPE_KW not in row["type"].lower(): - continue - if row["meter"].strip() != METER: - continue + if not a.mixed: + if a.type_kw not in row["type"].lower(): + continue + if row["meter"].strip() != a.meter: + continue body = clean_abc(row["abc"]) if not (40 <= len(body) <= 700): continue if any(ch not in ALLOWED for ch in body.replace("\n", "")): continue key = norm_key(row["mode"]) - block = f"X:1\nM:{METER}\nK:{key}\n{body}\n" + meter = row["meter"].strip() if a.mixed else a.meter + block = f"X:1\nM:{meter}\nK:{key}\n{body}\n" rows_out.append(block) n_kept += 1 text = "\n".join(rows_out) - os.makedirs(os.path.dirname(OUT) or ".", exist_ok=True) - with open(OUT, "w", encoding="utf-8") as f: + os.makedirs(os.path.dirname(a.out) or ".", exist_ok=True) + with open(a.out, "w", encoding="utf-8") as f: f.write(text) vocab = sorted(set(text)) sys.stdout.reconfigure(encoding="utf-8") print(f"melodii w pliku : {n_total}") - print(f"{TYPE_KW} ({METER}) zachowane: {n_kept}") + desc = "MIXED (bez filtra)" if a.mixed else f"{a.type_kw} ({a.meter})" + print(f"{desc} zachowane : {n_kept}") print(f"znaki łącznie : {len(text):,}") print(f"słownik ({len(vocab)}) : {''.join(vocab)!r}") print("\n--- pierwszy blok ---") diff --git a/music-experts/src/generate/make_midi.py b/music-experts/src/generate/make_midi.py index 717a9f1..1e5cff0 100644 --- a/music-experts/src/generate/make_midi.py +++ b/music-experts/src/generate/make_midi.py @@ -24,6 +24,7 @@ def first_tune(raw: str) -> str: def main(): ap = argparse.ArgumentParser() + ap.add_argument("--meter", default="6/8", help="metrum, np. 6/8, 4/4, 3/4") ap.add_argument("--key", default="D", help="tonacja, np. D, G, Am, Edor") ap.add_argument("--n", type=int, default=3, help="ile melodii") ap.add_argument("--temp", type=float, default=0.8, help="temperatura (więcej = śmielej)") @@ -42,7 +43,7 @@ def main(): model.eval() print(f"GPT {model.num_params():,} param | val loss {ck['val_loss']:.3f} | {device}") - seed = f"X:1\nM:6/8\nK:{args.key}\n" + seed = f"X:1\nM:{args.meter}\nK:{args.key}\n" made = 0 for i in range(1, args.n + 1): idx = torch.tensor([[stoi[c] for c in seed]], dtype=torch.long, device=device) diff --git a/music-experts/src/tools/benchmark_judge.py b/music-experts/src/tools/benchmark_judge.py new file mode 100644 index 0000000..84bdd11 --- /dev/null +++ b/music-experts/src/tools/benchmark_judge.py @@ -0,0 +1,231 @@ +# Autor: Adam Skrodzki +"""Benchmark eksperta GPT oceniany przez sędziego (judge). Dwa zadania: + + scratch — tylko etykiety (meter/mode) w promptcie `X:1\nM:m\nK:k\n`; model generuje + od zera; sędzia klasyfikuje ciało. + continuation — jak wyżej + w promptcie ćwiartka prawdziwego ciała melodii z val; model + kontynuuje; sędzia klasyfikuje CAŁOŚĆ (prefiks + kontynuacja) — pytanie + brzmi: czy wynikowy utwór nadal jest "w stylu". + +Domena: pula melodii = val split ograniczony do DOMENY checkpointu (data/models/domains.json: +type substring + meter, jak w prepare_data.py). Melodie spoza domeny nie są błędem modelu — +nie trafiają do puli wcale (bach: type=null = poza benchmarkiem). W domenie liczą się WSZYSTKIE +próbki: porażka modelu (śmieci zamiast melodii) = 0 punktu. + +Wynik: dla każdej próbki sumujemy prawdopodobieństwa, jakie sędzia przyznaje PRAWDZIWYM +klasom, uśredniamy -> score 0..1 (+ pokrycie). Dodatkowo accuracy (argmax == prawda) per +głowa. Pula = val sędziego (split po tune_id, ten sam seed co train_judge), więc sędzia +ocenia melodie, których nigdy nie widział. + +Użycie: + python src/tools/benchmark_judge.py --model data/models/jig_ckpt.pt + python src/tools/benchmark_judge.py --model data/models/jig_ckpt.pt --task continuation + python src/tools/benchmark_judge.py --model data/models/jig_ckpt.pt --json +""" +import argparse, collections, json, sys, os, random +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import torch +import torch.nn.functional as F +from core.gpt import GPT +from core.abc_corpus import parse_tune_row +from core.judge import JudgeGPT, HEADS +from train_judge import load_rows, group_split, MIN_BODY # reuse: czyszczenie + split +from judge_tunes import windows # reuse: okna ciała + +def mean_probs(judge, jcfg, jstoi, device, body, bs=64): + """Średni softmax sędziego po wszystkich oknach ciała -> {head: wektor prawd.}.""" + wins = windows(body, jcfg.block_size) + idx = torch.zeros(len(wins), jcfg.block_size, dtype=torch.long, device=device) + mask = torch.zeros(len(wins), jcfg.block_size, device=device) + for b, w in enumerate(wins): + ids = [jstoi[c] for c in w if c in jstoi] or [jstoi[" "]] + idx[b, :len(ids)] = torch.tensor(ids, device=device) + mask[b, :len(ids)] = 1.0 + with torch.no_grad(): + logits = judge(idx, mask) + return {h: F.softmax(logits[h], dim=-1).mean(dim=0) for h in HEADS} + +def run_task(model, stoi, itos, judge, jcfg, jstoi, classes, device, task, + rows, sample_ids, a, failed=None): + """Jedno zadanie benchmarku. Zwraca (score 0..1, statystyki, liczba pominiętych). + Scratch ocenia tylko meter i mode — type NIE jest w promptcie, model nie ma jak go + zgadnąć (generuje styl swojego korpusu), więc ocenianie go byłoby mylące. + Pominięte próbki (znaki spoza słownika, ciało < MIN_BODY) liczą się jako 0 — inaczej + survivor bias zawyżałby wynik słabych modeli (score tylko nad udanymi generacjami). + W continuation liczy też statystyki PRAWDZIWEJ melodii (całe ciało z val) — referencja; + referencja NIE zależy od modelu, więc liczona dla WSZYSTKICH próbek.""" + truth = {h: {c: i for i, c in enumerate(classes[h])} for h in HEADS} + heads = HEADS if task == "continuation" else [h for h in HEADS if h != "type"] + p_sum = {h: 0.0 for h in heads}; acc = {h: 0 for h in heads} + p_ref = {h: 0.0 for h in heads}; acc_ref = {h: 0 for h in heads} # referencja = prawdziwe ciało + score_sum, score_ref, skipped = 0.0, 0.0, 0 + skip_reasons = collections.Counter() + + def tally(row, probs): + """Zbierz (sum prawd. prawdziwych klas / liczba ocenianych głów, per-head p/acc).""" + ps = {h: 0.0 for h in heads}; ac = {h: 0 for h in heads} + for h in heads: + t = truth[h][row[h]] + ps[h] = float(probs[h][t]); ac[h] = int(int(probs[h].argmax()) == t) + return sum(ps.values()) / len(heads), ps, ac + + for s in sample_ids: + row = rows[s] + if task == "continuation": + # referencja najpierw: nie zależy od modelu, liczona nawet gdy model zawiedzie + s_true, ps, ac = tally(row, mean_probs(judge, jcfg, jstoi, device, row["body"])) + score_ref += s_true; p_ref.update((h, p_ref[h] + ps[h]) for h in heads) + acc_ref.update((h, acc_ref[h] + ac[h]) for h in heads) + if task == "scratch": + prompt = f"X:1\nM:{row['meter']}\nK:{row['mode']}\n" + else: # continuation: ćwiartka prawdziwego ciała jako kontekst stylu + body_true = row["body"] + prompt = f"X:1\nM:{row['meter']}\nK:{row['mode']}\n{body_true[:max(40, len(body_true) // 4)]}" + if any(c not in stoi for c in prompt): + skipped += 1 # = 0 punktu + missing = sorted({c for c in prompt if c not in stoi}) + skip_reasons["vocab"] += 1 + if failed is not None: + failed.append({"task": task, "reason": "vocab", "prompt": prompt, + "missing_chars": missing}) + continue + idx = torch.tensor([[stoi[c] for c in prompt]], dtype=torch.long, device=device) + with torch.no_grad(): + gen = model.generate(idx, a.new, temperature=a.temp, top_k=a.topk)[0].tolist() + raw = "".join(itos[t] for t in gen) + body = "\n".join(l for l in (prompt + raw).split("\n") + if not l.startswith(("X:", "T:", "C:", "M:", "K:", "N:", "%"))).strip() + if len(body) < MIN_BODY: + skipped += 1 # = 0 punktu + skip_reasons["short_body"] += 1 + if failed is not None: + failed.append({"task": task, "reason": "short_body", "prompt": prompt, + "generated": raw, "body_len": len(body)}) + continue + s_model, ps, ac = tally(row, mean_probs(judge, jcfg, jstoi, device, body)) + score_sum += s_model; p_sum.update((h, p_sum[h] + ps[h]) for h in heads) + acc.update((h, acc[h] + ac[h]) for h in heads) + n = max(len(sample_ids), 1) # mianownik = WSZYSTKIE próbki, pominięte wliczone jako 0 + stats = {h: {"p_correct": p_sum[h] / n, "acc": acc[h] / n} for h in heads} + if task == "continuation": + stats["actual"] = {"score": score_ref / n, + **{h: {"p_correct": p_ref[h] / n, "acc": acc_ref[h] / n} for h in heads}} + return score_sum / n, stats, skipped, dict(skip_reasons) + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--model", required=True, help="ckpt eksperta GPT (ten, którego oceniamy)") + ap.add_argument("--judge", default="data/models/judge_v2.pt") + ap.add_argument("--csv", default="data/tunes.csv") + ap.add_argument("--domains", default="data/models/domains.json", + help="mapowanie checkpoint -> domena (type/meter); pula = val w domenie") + ap.add_argument("--samples", type=int, default=100) + ap.add_argument("--task", default="both", choices=["scratch", "continuation", "both"]) + ap.add_argument("--seed", type=int, default=42, + help="seed losowania próbek i generacji; ten sam seed = IDENTYCZNE " + "melodie dla każdego modelu (porównywalne benchmarki)") + ap.add_argument("--split-seed", type=int, default=42, + help="seed splitu val — musi zgadzać się z seedem treningu sędziego") + ap.add_argument("--new", type=int, default=420, help="ile znaków generować") + ap.add_argument("--temp", type=float, default=0.85) + ap.add_argument("--topk", type=int, default=18) + ap.add_argument("--json", action="store_true") + ap.add_argument("--dump-failed", default=None, metavar="PLIK", + help="zapisz pominięte generacje (prompt + surowy output) do PLIKU (json)") + a = ap.parse_args() + sys.stdout.reconfigure(encoding="utf-8") + with open(a.domains, encoding="utf-8") as f: + domains = json.load(f) + device = "cuda" if torch.cuda.is_available() else "cpu" + torch.manual_seed(a.seed) # deterministyczna generacja + rng = random.Random(a.seed) # deterministyczny dobór próbek + print(f"urządzenie: {device} | seed: {a.seed}") + print(f"sędzia: {a.judge}") + ck = torch.load(a.judge, map_location="cpu", weights_only=False) + jcfg, chars, classes = ck["config"], ck["chars"], ck["classes"] + jstoi = {c: j + 1 for j, c in enumerate(chars)} + judge = JudgeGPT(jcfg, {h: len(classes[h]) for h in HEADS}, pool=ck.get("pool", "mean")) + judge.load_state_dict(ck["model"]); judge.eval().to(device) + judge_accs = ck.get("accs", {}) + + print(f"model: {a.model}") + mck = torch.load(a.model, map_location=device, weights_only=False) + stoi, itos, mcfg = mck["stoi"], mck["itos"], mck["config"] + model = GPT(mcfg); model.load_state_dict(mck["model"]); model.eval().to(device) + + print(f"czytam {a.csv} ...") + rows = load_rows(a.csv) + _, va = group_split(rows, a.split_seed) + + # Domena modelu (data/models/domains.json): benchmarkujemy TYLKO melodie z domeny + # checkpointu. Melodia spoza domeny nie jest błędem modelu — nie może jej być ani + # brnąć w nią; nie liczy się w ogóle. Błędy W domenie (śmieci, brak ciała) = 0. + dom = domains.get(os.path.basename(a.model)) + if dom is None: + sys.exit(f"brak domeny dla {os.path.basename(a.model)} w {a.domains} — dopisz wpis") + if dom.get("type") is None: + sys.exit(f"{os.path.basename(a.model)}: domena poza zakresem tunes.csv " + f"(type=null w {a.domains}) — benchmark nie ma czego mierzyć") + dom_pool = [i for i in va + if dom["type"] in rows[i]["type"].lower() + and rows[i]["meter"].strip() == dom["meter"]] + pool = [i for i in dom_pool if len(rows[i]["body"]) >= 4 * 40] # ćwiartka prefiksu >= 40 znaków + print(f"domena modelu: type~'{dom['type']}' + meter {dom['meter']}") + print(f"pula melodii: val {len(va)} -> w domenie {len(dom_pool)} -> z ciałem {len(pool)}") + if not pool: + sys.exit("pusta pula domenowa — benchmark niemożliwy") + print(f"val acc sędziego: " + " | ".join(f"{h} {v:.3f}" for h, v in judge_accs.items())) + sample_ids = rng.sample(pool, min(a.samples, len(pool))) + + tasks = ["scratch", "continuation"] if a.task == "both" else [a.task] + failed_all = [] + results = {} + for task in tasks: + score, stats, skipped, skip_reasons = run_task(model, stoi, itos, judge, jcfg, jstoi, + classes, device, task, rows, sample_ids, + a, failed=failed_all) + results[task] = {"score": score, "skipped": skipped, "skip_reasons": skip_reasons, **stats} + if a.json: + continue + ref = stats.get("actual") + why = ", ".join(f"{k}: {v}" for k, v in skip_reasons.items()) or "-" + cov = 100 * (len(sample_ids) - skipped) / max(len(sample_ids), 1) + print(f"\n=== {task} (score {score:.3f} @ pokrycie {cov:.0f}%" + + (f", prawdziwa melodia {ref['score']:.3f}" if ref else "") + + f", pominięto {skipped}/{len(sample_ids)} [{why}]) ===") + for h in [h for h in HEADS if h in stats]: + st = stats[h] + line = f" {h:6s}: P(prawdziwa klasa) {st['p_correct']:.3f} | acc {st['acc']:.3f}" + if ref: + ra = ref[h] + line += f" [prawdziwa: P {ra['p_correct']:.3f} | acc {ra['acc']:.3f}]" + print(line) + + if a.dump_failed and failed_all: + os.makedirs(os.path.dirname(a.dump_failed) or ".", exist_ok=True) + with open(a.dump_failed, "w", encoding="utf-8") as f: + json.dump({"model": a.model, "reasons": dict(collections.Counter( + r["reason"] for r in failed_all)), "failed": failed_all}, + f, ensure_ascii=False, indent=2) + print(f"\npominięte generacje ({len(failed_all)}) -> {a.dump_failed}") + + if a.json: + print(json.dumps({"model": a.model, "judge": a.judge, "samples": a.samples, + "judge_val_accs": judge_accs, "tasks": results}, + ensure_ascii=False, indent=2)) + return + + print(f"\nbenchmark: {a.model} | {len(sample_ids)} melodii z val | temp {a.temp} topk {a.topk}") + for task in tasks: + r = results[task] + cov = 100 * (len(sample_ids) - r["skipped"]) / max(len(sample_ids), 1) + line = f" {task:12s} score: {r['score']:.3f} @ pokrycie {cov:.0f}%" + if "actual" in r: + line += f" (prawdziwe melodie: {r['actual']['score']:.3f})" + print(line) + print("score = średnie prawdopodobieństwo sędziego przy prawdziwych klasach (0..1); " + "pominięte generacje liczą się jako 0. Score porównuj tylko przy zbliżonym pokryciu — " + "niskie pokrycie = model nie jest w stanie nawet zakodować promptu (inny alfabet korpusu)") + +if __name__ == "__main__": + main() diff --git a/music-experts/src/tools/benchmark_ood.py b/music-experts/src/tools/benchmark_ood.py new file mode 100644 index 0000000..d49d738 --- /dev/null +++ b/music-experts/src/tools/benchmark_ood.py @@ -0,0 +1,184 @@ +# Autor: Adam Skrodzki +"""OOD transfer matrix: ekspert × docelowa domena. Jak wypada model POZA swoją domeną? + +Dla każdej pary (model, cel): pula = val melodie DOMENY CELU (type+meter jak prepare_data), +prompt z nagłówkami celu (scratch) lub z ćwiartką ciała melodii celu (continuation). +PRAWDA = to, o co prosimy (etykiety z melodii celu), nie domena modelu — model jest proszony, +więc nie ma wymówki. Porażka w OOD (śmieci, znaki spoza słownika) = 0, jak w benchmarku +domenowym; pokrycie raportowane osobno. + +Metryki w komórce macierzy: score/ref — score = średnie P sędziego przy prawdziwych klasach, +ref = to samo dla PRAWDZIWYCH melodii celu (sufit). hb = home-bias: P sędziego przy KLASIE +DOMOWEJ modelu (meter dla scratch; średnia meter+type dla continuation) — rozróżnia sztywnego +modelu (wysoki hb, ignoruje prompt) od papki (niskie wszystko). * = komórka domenowa (diagonala). + +Użycie: + python src/tools/benchmark_ood.py # wszystkie ckpt z domenami + python src/tools/benchmark_ood.py --models data/models/jig_ckpt.pt --samples 50 + python src/tools/benchmark_ood.py --json > data/benchmarks/ood.json +""" +import argparse, glob, json, sys, os, random +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import torch +from core.gpt import GPT +from core.judge import JudgeGPT, HEADS +from train_judge import load_rows, group_split, MIN_BODY +from benchmark_judge import mean_probs + +def gen_body(model, stoi, itos, device, prompt, a): + """Ciało wygenerowanej melodii (bez linii nagłówkowych) albo None, gdy prompt + niekodowalny / ciało zbyt krótkie (porażka OOD = 0 punktu).""" + if any(c not in stoi for c in prompt): + return None, "vocab" + idx = torch.tensor([[stoi[c] for c in prompt]], dtype=torch.long, device=device) + with torch.no_grad(): + gen = model.generate(idx, a.new, temperature=a.temp, top_k=a.topk)[0].tolist() + body = "\n".join(l for l in (prompt + "".join(itos[t] for t in gen)).split("\n") + if not l.startswith(("X:", "T:", "C:", "M:", "K:", "N:", "%"))).strip() + return (body, None) if len(body) >= MIN_BODY else (None, "short_body") + +def run_cell(model, stoi, itos, judge, jcfg, jstoi, classes, device, rows, + sample_ids, home, task, a): + """Jedna komórka macierzy (model × cel × zadanie). Zwraca statystyki.""" + truth = {h: {c: i for i, c in enumerate(classes[h])} for h in HEADS} + heads = HEADS if task == "continuation" else [h for h in HEADS if h != "type"] + home_heads = [h for h in ("meter", "type") if h in heads and home.get(h) in truth[h]] + p_sum = {h: 0.0 for h in heads}; hb_sum = {h: 0.0 for h in home_heads} + score, invalid = 0.0, {"vocab": 0, "short_body": 0} + for i in sample_ids: + row = rows[i] + prompt = f"X:1\nM:{row['meter']}\nK:{row['mode']}\n" + if task == "continuation": + prompt += row["body"][:max(40, len(row["body"]) // 4)] + body, why = gen_body(model, stoi, itos, device, prompt, a) + if body is None: + invalid[why] += 1 + continue # = 0 punktu + probs = mean_probs(judge, jcfg, jstoi, device, body) + for h in heads: + t = truth[h][row[h]] + score += float(probs[h][t]) + p_sum[h] += float(probs[h][t]) + if h in home_heads: + hb_sum[h] += float(probs[h][truth[h][home[h]]]) + n = max(len(sample_ids), 1) + nh = len(heads) + return {"score": score / (n * nh), + "p": {h: p_sum[h] / n for h in heads}, + "home_bias": {h: hb_sum[h] / n for h in home_heads}, + "invalid": invalid, "coverage": 1 - sum(invalid.values()) / n} + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--models", nargs="*", default=None, + help="ckpty do macierzy (domyślnie wszystkie *_ckpt.pt z domeną)") + ap.add_argument("--judge", default="data/models/judge_v2.pt") + ap.add_argument("--csv", default="data/tunes.csv") + ap.add_argument("--domains", default="data/models/domains.json") + ap.add_argument("--samples", type=int, default=100, help="melodii na komórkę") + ap.add_argument("--seed", type=int, default=42) + ap.add_argument("--split-seed", type=int, default=42) + ap.add_argument("--new", type=int, default=420) + ap.add_argument("--temp", type=float, default=0.85) + ap.add_argument("--topk", type=int, default=18) + ap.add_argument("--json", action="store_true") + a = ap.parse_args() + sys.stdout.reconfigure(encoding="utf-8") + device = "cuda" if torch.cuda.is_available() else "cpu" + torch.manual_seed(a.seed) + print(f"urządzenie: {device} | seed: {a.seed} | temp {a.temp} topk {a.topk} new {a.new}") + + with open(a.domains, encoding="utf-8") as f: + domains = json.load(f) + + ck = torch.load(a.judge, map_location="cpu", weights_only=False) + jcfg, chars, classes = ck["config"], ck["chars"], ck["classes"] + jstoi = {c: j + 1 for j, c in enumerate(chars)} + judge = JudgeGPT(jcfg, {h: len(classes[h]) for h in HEADS}, pool=ck.get("pool", "mean")) + judge.load_state_dict(ck["model"]); judge.eval().to(device) + + # cele = wszystkie różne domeny (type, meter); diag = domena modelu + # (_comment w domains.json to string — bierzemy tylko wpisy słownikowe) + targets = sorted({(d["type"], d["meter"]) for n, d in domains.items() + if isinstance(d, dict) and d.get("type") and not n.startswith("judge")}) + print(f"cele: " + ", ".join(f"{t}:{m}" for t, m in targets)) + + print(f"czytam {a.csv} ...") + rows = load_rows(a.csv) + _, va = group_split(rows, a.split_seed) + + # pula i próbki per CEL: niezależne od modeli (ten sam seed -> te same melodie + # w kolumnie dla każdego modelu) + referencja = prawdziwe melodie celu + pools, refs = {}, {} + for ttype, tmeter in targets: + key = f"{ttype}:{tmeter}" + rng = random.Random(f"{a.seed}|{key}") # deterministyczny per cel + pool = [i for i in va + if ttype in rows[i]["type"].lower() + and rows[i]["meter"].strip() == tmeter + and len(rows[i]["body"]) >= 4 * 40] + ids = rng.sample(pool, min(a.samples, len(pool))) + pools[key] = ids + # referencja: sędzia na prawdziwych ciałach celu (nie zależy od modelu) + s = {h: 0.0 for h in HEADS} + for i in ids: + probs = mean_probs(judge, jcfg, jstoi, device, rows[i]["body"]) + for h in HEADS: + s[h] += float(probs[h][classes[h].index(rows[i][h])]) + refs[key] = sum(s.values()) / (len(ids) * len(HEADS)) + print(f" cel {key}: pula {len(pool)}, referencja {refs[key]:.3f}") + + models = a.models or sorted(glob.glob("data/models/*_ckpt.pt")) + cells = {} + for mp in models: + name = os.path.basename(mp) + dom = domains.get(name) + if dom is None or dom.get("type") is None: + print(f"pomijam {name}: brak domeny w {a.domains}") + continue + home = {"type": dom["type"], "meter": dom["meter"]} + mck = torch.load(mp, map_location=device, weights_only=False) + stoi, itos, mcfg = mck["stoi"], mck["itos"], mck["config"] + model = GPT(mcfg); model.load_state_dict(mck["model"]); model.eval().to(device) + for ttype, tmeter in targets: + key = f"{ttype}:{tmeter}" + diag = (ttype == dom["type"] and tmeter == dom["meter"]) + for task in ("scratch", "continuation"): + st = run_cell(model, stoi, itos, judge, jcfg, jstoi, classes, device, + rows, pools[key], home, task, a) + cells[(name, key, task)] = st + if not a.json: + print(f" {name} x {key} [{task}]: score {st['score']:.3f} " + f"(hb {st['home_bias']}, pokrycie {st['coverage']:.0%})") + del model + torch.cuda.empty_cache() if device == "cuda" else None + + if a.json: + print(json.dumps({"seed": a.seed, "targets": [f"{t}:{m}" for t, m in targets], + "refs": refs, "cells": {f"{n}|{k}|{t}": v + for (n, k, t), v in cells.items()}}, + ensure_ascii=False, indent=2)) + return + + # macierze: wiersz = model, kolumna = cel; komórka score/ref hb.. c..% (* = diag) + for task in ("scratch", "continuation"): + w = 24 + print(f"\n=== {task} | komórka: score/ref h=home-bias c=pokrycie ===") + print(f"{'model':22s}" + "".join(f"{k:>{w}}" for k in refs)) + for name in sorted({n for n, _, _ in cells}): + row = f"{name:22s}" + for key in refs: + st = cells.get((name, key, task)) + if st is None: + row += "—".rjust(w) + continue + diag = "*" if domains.get(name, {}).get("type") == key.split(":")[0] else "" + hb = "/".join(f"{v:.2f}" for v in st["home_bias"].values()) or "—" + cell = f"{st['score']:.2f}/{refs[key]:.2f}h{hb}c{st['coverage']:.0%}{diag}" + row += cell[:w].rjust(w) + print(row) + print("\nref = sędzia na prawdziwych melodiach celu (sufit); hb wysokie = model ignoruje " + "prompt i wraca do domeny; niskie score i niskie hb = papka. * = komórka domenowa.") + +if __name__ == "__main__": + main() diff --git a/music-experts/src/tools/judge_tunes.py b/music-experts/src/tools/judge_tunes.py new file mode 100644 index 0000000..404dcaf --- /dev/null +++ b/music-experts/src/tools/judge_tunes.py @@ -0,0 +1,94 @@ +# Autor: Adam Skrodzki +"""Judge: klasyfikuje wygenerowane melodie ABC (meter/mode/type) z samego ciała. +Trunk GPT + mean-pool (core/judge.py). Ciało dzielone na okna 256 znaków, +softmax uśredniony po oknach (pełne pokrycie — jak w evaluate() train_judge.py). +Wejście: pliki .abc lub katalogi. Wynik: tabela + agregat; --json dla maszynowej postaci. +Użycie: python src/tools/judge_tunes.py out/ [--json] [--judge data/models/judge_v2.pt] +""" +import argparse, glob, json, os, sys +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +import torch +import torch.nn.functional as F +from core.gpt import GPTConfig +from core.judge import JudgeGPT, HEADS +from core.abc_corpus import tune_body + +def split_tunes(text): + """Kolejne bloki melodii z pliku ABC (blok kończy się na następnej linii X:).""" + tunes, cur = [], [] + for ln in text.split("\n"): + if ln.startswith("X:") and cur: + tunes.append("\n".join(cur)); cur = [] + cur.append(ln) + if cur and any(l.startswith("X:") for l in cur): + tunes.append("\n".join(cur)) + return tunes + +def windows(body, block): + return [body[s:s + block] for s in range(0, len(body), block)] or [" "] + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("inputs", nargs="+", help="pliki .abc lub katalogi") + ap.add_argument("--judge", default="data/models/judge_v2.pt") + ap.add_argument("--json", action="store_true") + a = ap.parse_args() + sys.stdout.reconfigure(encoding="utf-8") + + ck = torch.load(a.judge, map_location="cpu", weights_only=False) + cfg, chars, classes = ck["config"], ck["chars"], ck["classes"] + stoi = {c: j + 1 for j, c in enumerate(chars)} + device = "cuda" if torch.cuda.is_available() else "cpu" + model = JudgeGPT(cfg, {h: len(classes[h]) for h in HEADS}, pool=ck.get("pool", "mean")) + model.load_state_dict(ck["model"]); model.eval().to(device) + print(f"urządzenie: {device}") + + files = [] + for path in a.inputs: + if os.path.isdir(path): + files += sorted(glob.glob(os.path.join(path, "**", "*.abc"), recursive=True)) + else: + files.append(path) + if not files: + sys.exit("brak plików .abc do oceny") + + tunes = [] + for path in files: + for ti, tune in enumerate(split_tunes(open(path, encoding="utf-8").read())): + body = tune_body(tune) + if len(body) >= 40: + tunes.append({"file": path, "tune": ti, "body": body}) + + results = [] + with torch.no_grad(): + for t in tunes: + wins = windows(t["body"], cfg.block_size) + idx = torch.zeros(len(wins), cfg.block_size, dtype=torch.long, device=device) + mask = torch.zeros(len(wins), cfg.block_size, device=device) + for b, w in enumerate(wins): + ids = [stoi[c] for c in w if c in stoi] or [stoi[" "]] + idx[b, :len(ids)] = torch.tensor(ids, device=device) + mask[b, :len(ids)] = 1.0 + logits = model(idx, mask) + rec = {"file": t["file"], "tune": t["tune"], "windows": len(wins)} + for h in HEADS: + mean_p = F.softmax(logits[h], dim=-1).mean(dim=0) + best = int(mean_p.argmax()) + rec[h] = classes[h][best] + rec[f"{h}_conf"] = round(float(mean_p[best]), 3) + results.append(rec) + + if a.json: + print(json.dumps(results, ensure_ascii=False, indent=2)) + return + + print(f"{'plik':40s} {'okna':4s} {'meter':7s} {'mode':7s} {'type':10s} {'pewność (m/m/t)'}") + for r in results: + conf = f"{r['meter_conf']:.2f}/{r['mode_conf']:.2f}/{r['type_conf']:.2f}" + print(f"{r['file']:40s} {r['windows']:4d} {r['meter']:7s} {r['mode']:7s} {r['type']:10s} {conf}") + accs = ck.get("accs", {}) + print(f"\noceniono: {len(results)} melodii z {len(files)} plików | judge: {a.judge} " + + ("(val acc: " + ", ".join(f"{h} {v:.3f}" for h, v in accs.items()) + ")" if accs else "")) + +if __name__ == "__main__": + main() diff --git a/music-experts/src/tools/train_judge.py b/music-experts/src/tools/train_judge.py new file mode 100644 index 0000000..ce50a1d --- /dev/null +++ b/music-experts/src/tools/train_judge.py @@ -0,0 +1,204 @@ +# Autor: Adam Skrodzki +"""Judge v1: klasyfikator meter/mode/type na trunku GPT (verbatim core/gpt.py). +Ciało melodii (nagłówki usunięte) -> znaki -> GPT -> pooling (mean|attn) -> 3 liniowe głowy. +Split po tune_id (settingi tej samej melodii nie mogą przeciekać do val). +Czysty PyTorch. Opcjonalnie --init-from: trunk startuje z wag pretrenowanego eksperta. +Użycie: python src/tools/train_judge.py [--csv data/tunes.csv] [--out data/models/judge_v2.pt] + [--init-from data/models/jig_ckpt.pt] [--seed 42] +""" +import argparse, collections, csv, os, random, sys, time +from contextlib import nullcontext +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +csv.field_size_limit(10**7) +import torch +import torch.nn.functional as F +from core.gpt import GPTConfig +from core.judge import JudgeGPT, HEADS +from core.abc_corpus import parse_tune_row + +MIN_BODY, MAX_BODY = 40, 700 # jak prepare_data.py +BLOCK = 256 +DEVICE = "cpu" # ustawiane w main(): "cuda" gdy dostępne (autodetekcja jak train_gpt.py) + +def load_rows(csv_path): + rows = [] + with open(csv_path, encoding="utf-8", newline="") as f: + for row in csv.DictReader(f): + body, meter, mode, typ = parse_tune_row(row) + if MIN_BODY <= len(body) <= MAX_BODY: + rows.append({"body": body, "tune_id": row["tune_id"], + "meter": meter, "mode": mode, "type": typ}) + return rows + +def group_split(rows, seed): + """90/10 po unikalnych tune_id (deterministycznie).""" + ids = sorted({r["tune_id"] for r in rows}) + rng = random.Random(seed) + rng.shuffle(ids) + n_val = max(1, len(ids) // 10) + val_ids = set(ids[:n_val]) + tr = [i for i, r in enumerate(rows) if r["tune_id"] not in val_ids] + va = [i for i, r in enumerate(rows) if r["tune_id"] in val_ids] + return tr, va + +def encode_train(rows, idxs, stoi, block): + """Trening: losowe okno block znaków z każdego ciała (augmentacja — przez epoki + model widzi całe ciało, nie tylko pierwsze block znaków). Krótsze ciała = 1 okno. + Padding na końcu; 0 w idx = padding.""" + B = len(idxs) + idx = torch.zeros(B, block, dtype=torch.long, device=DEVICE) + mask = torch.zeros(B, block, device=DEVICE) + for b, i in enumerate(idxs): + body = rows[i]["body"] + start = random.randint(0, max(0, len(body) - block)) + ids = [stoi[c] for c in body[start:start + block] if c in stoi] + if not ids: + ids = [stoi[" "]] + idx[b, :len(ids)] = torch.tensor(ids, device=DEVICE) + mask[b, :len(ids)] = 1.0 + return idx, mask + +def windows(body, block): + """Wszystkie okna niepokrywające (ostatnie może być krótsze).""" + return [body[s:s + block] for s in range(0, len(body), block)] or [" "] + +def evaluate(model, rows, idxs, classes, stoi, bs=64, block=BLOCK): + """Eval: KAŻDE ciało dzielone na wszystkie okna (pełne pokrycie, deterministycznie), + softmax uśredniany po oknach -> jedna predykcja na ciało.""" + model.eval() + samples = [(i, w) for i in idxs for w in windows(rows[i]["body"], block)] + agg = {i: {h: torch.zeros(len(classes[h])) for h in HEADS} for i in idxs} + cnt = collections.Counter() + with torch.no_grad(): + for s in range(0, len(samples), bs): + chunk = samples[s:s+bs] + idx = torch.zeros(len(chunk), block, dtype=torch.long, device=DEVICE) + mask = torch.zeros(len(chunk), block, device=DEVICE) + for b, (i, w) in enumerate(chunk): + ids = [stoi[c] for c in w if c in stoi] or [stoi[" "]] + idx[b, :len(ids)] = torch.tensor(ids, device=DEVICE) + mask[b, :len(ids)] = 1.0 + logits = model(idx, mask) + for b, (i, _) in enumerate(chunk): + for h in HEADS: + agg[i][h] += F.softmax(logits[h][b], dim=-1).cpu() + cnt[i] += 1 + correct = collections.Counter(); n = 0 + confusion = {h: collections.Counter() for h in HEADS} + for i in idxs: + n += 1 + for h in HEADS: + cls = classes[h] + mean_p = agg[i][h] / cnt[i] + pred = cls[int(mean_p.argmax())] + true = rows[i][h] + correct[h] += pred == true + if pred != true: + confusion[h].update([(true, pred)]) + model.train() + return {h: correct[h] / n for h in HEADS}, confusion + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--csv", default="data/tunes.csv") + ap.add_argument("--out", default="data/models/judge_v2.pt") + ap.add_argument("--init-from", default=None, help="ckpt eksperta: pretrenowany trunk") + ap.add_argument("--seed", type=int, default=42) + ap.add_argument("--iters", type=int, default=2000) + ap.add_argument("--batch", type=int, default=32) + ap.add_argument("--lr", type=float, default=3e-4) + ap.add_argument("--block", type=int, default=256, help="długość okna ciała (znaki)") + ap.add_argument("--pool", choices=["mean", "attn"], default="mean", + help="pooling po pozycjach: mean (domyślny w dotychczasowych sędziach) albo attention-pool") + ap.add_argument("--eval-every", type=int, default=100) + ap.add_argument("--train-eval-n", type=int, default=5000, + help="ile przykładów train do logowania acc (0 = cały train)") + a = ap.parse_args() + sys.stdout.reconfigure(encoding="utf-8") + global DEVICE + DEVICE = "cuda" if torch.cuda.is_available() else "cpu" + use_bf16 = DEVICE == "cuda" and torch.cuda.is_bf16_supported() + ctx = torch.autocast(device_type="cuda", dtype=torch.bfloat16) if use_bf16 else nullcontext() + torch.manual_seed(a.seed) + print(f"urządzenie: {DEVICE} | bf16: {use_bf16}") + + print(f"czytam {a.csv} ...") + rows = load_rows(a.csv) + tr, va = group_split(rows, a.seed) + steps_per_epoch = len(tr) // a.batch + TRAIN_EVAL = tr if a.train_eval_n <= 0 else tr[:a.train_eval_n] # stała próbka train do logu + print(f"melodii (próbek): {len(rows)} łącznie | train {len(tr)} | val {len(va)} (split po tune_id)") + print(f"batch: {a.batch} próbek/krok | okno: {a.block} znaków | lr: {a.lr}") + print(f"próbkowanie losowe: pełne 'przejście' po train ≈ co {steps_per_epoch} kroków " + f"({len(tr)}/{a.batch}); ten sam przykład wraca średnio co ~{steps_per_epoch} kroków, " + f"ale okno (fragment ciała) losowane jest za każdym razem na nowo") + print(f"logowanie acc co {a.eval_every} kroków: train = stała próbka {len(TRAIN_EVAL)} " + f"przykładów, val = wszystkie {len(va)}") + + chars = sorted({c for i in tr for c in rows[i]["body"]}) + stoi = {c: j + 1 for j, c in enumerate(chars)} # 0 = padding + print(f"słownik ciał: {len(chars)} znaków (+padding)") + + classes = {h: sorted({r[h] for r in rows}) for h in HEADS} + n_classes = {h: len(c) for h, c in classes.items()} + print("klasy: " + " | ".join(f"{h}: {n_classes[h]}" for h in HEADS)) + + cfg = GPTConfig(vocab_size=len(chars) + 1, block_size=a.block) + model = JudgeGPT(cfg, n_classes, pool=a.pool).to(DEVICE) + + if a.init_from: + ck = torch.load(a.init_from, map_location="cpu", weights_only=False) + src = ck["model"] + dst = model.state_dict() + copied = [k for k in dst if k.startswith("gpt.") and k != "gpt.head.weight" + and k in src and src[k].shape == dst[k].shape] + for k in copied: + dst[k] = src[k] + model.load_state_dict(dst) + print(f"trunk zainicjowany z {a.init_from} ({len(copied)} tensorów, LM-head pominięty)") + print(f"parametry: {sum(p.numel() for p in model.parameters()):,}") + + opt = torch.optim.AdamW(model.parameters(), lr=a.lr, betas=(0.9, 0.99), weight_decay=0.1) + + def train_step(): + chunk = random.sample(tr, a.batch) + idx, mask = encode_train(rows, chunk, stoi, a.block) + with ctx: + logits = model(idx, mask) + loss = sum(F.cross_entropy(logits[h], torch.tensor([classes[h].index(rows[i][h]) for i in chunk], device=DEVICE)) + for h in HEADS) + opt.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + opt.step() + return loss.item() + + best_acc, t0 = -1.0, time.time() + for it in range(a.iters + 1): + if it % a.eval_every == 0 or it == a.iters: + tr_accs, _ = evaluate(model, rows, TRAIN_EVAL, classes, stoi, block=a.block) + va_accs, confusion = evaluate(model, rows, va, classes, stoi, block=a.block) + print(f"iter {it:4d} | train acc: " + + " | ".join(f"{h} {tr_accs[h]:.3f}" for h in HEADS) + + f" | val acc: " + + " | ".join(f"{h} {va_accs[h]:.3f}" for h in HEADS) + + f" | {time.time()-t0:.0f}s") + if it > 0 and sum(va_accs.values()) > best_acc: + best_acc = sum(va_accs.values()) + torch.save({"model": model.state_dict(), "config": cfg, "chars": chars, + "classes": classes, "accs": va_accs, "seed": a.seed, + "pool": model.pool}, a.out) + if it == a.iters: + break + train_step() + + accs, confusion = evaluate(model, rows, va, classes, stoi, block=a.block) + tr_accs, _ = evaluate(model, rows, TRAIN_EVAL, classes, stoi, block=a.block) + print("\n--- val (ostateczna, pełny val) vs train (stała próbka) ---") + for h in HEADS: + print(f"{h}: val acc {accs[h]:.4f} | train acc {tr_accs[h]:.4f} " + f"| top pomyłki: {confusion[h].most_common(5)}") + print(f"\nbest val (suma acc) zapisane -> {a.out}") + +if __name__ == "__main__": + main() diff --git a/music-experts/src/train/train_gpt.py b/music-experts/src/train/train_gpt.py index 9e53cb0..70c0cda 100644 --- a/music-experts/src/train/train_gpt.py +++ b/music-experts/src/train/train_gpt.py @@ -1,12 +1,23 @@ """Trening GPT od zera na korpusie ABC (L03/L08 w praktyce). Batche -> strata cross-entropy -> backprop -> AdamW -> val loss -> checkpoint. """ -import os, time, math, sys +import os, time, math, sys, argparse from contextlib import nullcontext import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from core.gpt import GPT, GPTConfig +ap = argparse.ArgumentParser() +ap.add_argument("data", nargs="?", default="data/jigs.abc") +ap.add_argument("ckpt", nargs="?", default="data/models/jig_ckpt.pt") +ap.add_argument("losslog", nargs="?", default="data/models/jig_loss_log.csv") +ap.add_argument("vocab_from", nargs="?", default=None, + help="wspólny słownik z innego ckpt (do stitchu)") +ap.add_argument("seed", nargs="?", type=int, default=20260620) # niezależne modele do E_CKA/E0.5 +ap.add_argument("--max-iters", type=int, default=2000, + help="skalować z rozmiarem korpusu (mixed ~4x jig)") +a = ap.parse_args() + # --- hiperparametry --- block_size = 128 batch_size = 32 @@ -15,12 +26,12 @@ n_embd = int(os.environ.get("N_EMBD", 128)) # ENV: sweep skali (E_CKA) dropout = 0.1 lr = 3e-4 -max_iters = 2000 +max_iters = a.max_iters eval_interval = 200 eval_iters = 100 warmup = 100 -SEED = int(sys.argv[5]) if len(sys.argv) > 5 else 20260620 # argv[5]: seed (niezależne modele do E_CKA/E0.5) +SEED = a.seed torch.manual_seed(SEED) device = "cuda" if torch.cuda.is_available() else "cpu" use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported() @@ -28,12 +39,9 @@ sys.stdout.reconfigure(encoding="utf-8") print(f"urządzenie: {device} | bf16: {use_bf16}") -# --- ścieżki z argumentów (domyślnie: jigi) --- -DATA = sys.argv[1] if len(sys.argv) > 1 else "data/jigs.abc" -CKPT = sys.argv[2] if len(sys.argv) > 2 else "data/models/jig_ckpt.pt" -LOSSLOG = sys.argv[3] if len(sys.argv) > 3 else "data/models/jig_loss_log.csv" -VOCAB_FROM = sys.argv[4] if len(sys.argv) > 4 else None # wspólny słownik z innego ckpt (do stitchu) -print(f"dane: {DATA} -> checkpoint: {CKPT}") +# --- ścieżki z argumentów --- +DATA, CKPT, LOSSLOG, VOCAB_FROM = a.data, a.ckpt, a.losslog, a.vocab_from +print(f"dane: {DATA} -> checkpoint: {CKPT} | max_iters: {max_iters}") # --- dane: char-level --- text = open(DATA, encoding="utf-8").read() diff --git a/music-experts/src/train/train_rkl.py b/music-experts/src/train/train_rkl.py new file mode 100644 index 0000000..3184af9 --- /dev/null +++ b/music-experts/src/train/train_rkl.py @@ -0,0 +1,297 @@ +# Autor: Adam Skrodzki +"""On-policy destylacja: KL student↔rotujący nauczyciel (continuation + scratch). + +Student generuje rollouty z promptów domeny-nauczyciela (temp/topk jak w benchmarku), +nauczyciel kompetentny w tej domenie ocenia wygenerowane znaki, a strata to per-token +KL między rozkładami studenta i nauczyciela — tylko na WYGENEROWANYCH pozycjach, liczone +w oknach 128, gdzie pierwsze 32 znaki okna to tylko kontekst (nauczyciel +i student przewidują z prawdziwego prefixu, nie z amputowanego środka melodii). +Nauczyciel nigdy nie widzi surowych danych innej domeny niż jego własna — student +uczy się „z drugiej ręki", z samych prawdopodobieństw. + +α-mixing: p_mix = (1−α)·p_teacher + α/|V|, |V| = rozmiar wspólnego słownika (54). +Podłoga prawdopodobieństwa α/|V| na każdym znaku — kalibracja na niewytrenowanym +supportie nauczyciela (znaki spoza jego korpusu mają losowe logity po weight-tyingu) ++ ograniczenie gradientów tam, gdzie nauczyciel jest pewny, a student generuje śmieci. +Koszt: lekkie ściągnięcie ku uniform — entropia generacji jest metryką kontrolną. + +Wspólny słownik: student i wszyscy nauczyciele muszą mieć identyczny zbiór znaków +(assert). Rotacja PER BATCH: domena (równo jig→reel→waltz) × zadanie (continuation: +prompt o STAŁEJ długości 96 znaków = nagłówki + prefiks ciała; scratch: same nagłówki, +batch z jednego kubełka długości nagłówka). Ewaluacja co eval_interval: KL / entropia / +log-prob nauczyciela na utwalonej puli val per domena i zadanie. Zapisywane oba +checkpointy: best (min val KL — słabo skorelowany z jakością benchmarkową) i last. + +Użycie: + python src/train/train_rkl.py \ + --student data/models/universalist_ckpt.pt \ + --teachers jig=data/models/jig_sh_ckpt.pt reel=data/models/reel_sh_ckpt.pt \ + waltz=data/models/waltz_sh_ckpt.pt \ + --max-iters 4000 +""" +import argparse, json, math, os, random, sys, time +from contextlib import nullcontext +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) # src/ +sys.path.insert(0, os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "tools")) # src/tools +import torch +import torch.nn.functional as F +from core.gpt import GPT +from train_judge import load_rows, group_split + +MIN_POOL_BODY = 4 * 40 # jak benchmark OOD: ćwiartka prefiksu >= 40 znaków + + +def build_prompt(row, task, total): + """continuation: nagłówki + prefiks ciała o STAŁEJ długości total (równa długość + pozwala generować rollouty batchowo — generate() nie maskuje paddingu). + scratch: same nagłówki; długość = długość nagłówków (batch z jednego kubełka długości).""" + head = f"X:1\nM:{row['meter']}\nK:{row['mode']}\n" + if task == "scratch": + return head + return (head + row["body"][:max(1, total - len(head))])[:total] + + +def window_loss(seq, plen, student, teacher, t2s, alpha, direction, ctx, device, ctx_chars=32): + """KL per-token na wygenerowanych pozycjach sekwencji (B, L). Okna ze stride: pierwsze + ctx_chars znaków okna (poza pierwszym) to tylko KONTEKST — nauczyciel i student + przewidują z prawdziwego prefixu. + Zwraca (loss, mean_log_p_mix_na_sample, mean_entropia_studenta, n_pozycji). + Statystyki per POZYCJA per SEKWENCJA (licznik n uwzględnia batch).""" + B = seq.size(0) + BLOCK = min(student.cfg.block_size, teacher.cfg.block_size) + V = student.cfg.vocab_size + stride = BLOCK - ctx_chars + loss_sum, tlp_sum, ent_sum, n = 0.0, 0.0, 0.0, 0 + for s in range(0, max(1, seq.size(1) - 1), stride): + w = seq[:, s:s + BLOCK] + T = w.size(1) + if T < 2: + continue + lo = ctx_chars if s > 0 else 0 # w 1. oknie kontekstem jest sam prompt + p = torch.arange(lo, T - 1, device=device) # lokalne pozycje logitów (predykcja p+1) + g = s + p + 1 # globalne indeksy znaków targetowych + m = g >= plen # tylko wygenerowane znaki + if not bool(m.any()): + continue + with torch.no_grad(), ctx: + tl, _ = teacher(w) + tp = F.softmax(tl.float(), dim=-1)[:, :, t2s] # mapowanie na porządek studenta + pmix = (1 - alpha) * tp + alpha / V + log_pmix = pmix.clamp_min(1e-12).log() + with ctx: + sl, _ = student(w) + log_ps = sl.float().log_softmax(-1) + ps = log_ps.exp() + if direction == "reverse": + lp = (ps * (log_ps - log_pmix)).sum(-1) # KL(student || mix) + else: + lp = (pmix * (log_pmix - log_ps)).sum(-1) # KL(mix || student) + tgt = seq[:, s + 1:s + T] + lp_m = lp[:, :-1][:, lo:][:, m] + tlp_m = log_pmix[:, :-1].gather(-1, tgt.unsqueeze(-1))[:, :, 0][:, lo:][:, m] + ent_m = (-(ps * log_ps).sum(-1))[:, :-1][:, lo:][:, m] + loss_sum = loss_sum + lp_m.sum() + tlp_sum = tlp_sum + tlp_m.sum() + ent_sum = ent_sum + ent_m.sum() + n = n + m.sum() * B # maska wspólna dla całego batcha + n = max(n, 1) + return loss_sum / n, tlp_sum / n, ent_sum / n, n + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--student", default="data/models/universalist_ckpt.pt", + help="checkpoint startowy studenta (musi mieć wspólny słownik z nauczycielami)") + ap.add_argument("--teachers", nargs="+", required=True, + help="domena=ckpt, np. jig=data/models/jig_sh_ckpt.pt") + ap.add_argument("--domains", default="data/models/domains.json") + ap.add_argument("--csv", default="data/tunes.csv") + ap.add_argument("--alpha", type=float, default=0.1) + ap.add_argument("--direction", choices=["reverse", "forward"], default="reverse") + ap.add_argument("--tasks", choices=["continuation", "scratch", "both"], default="both", + help="zadania w rotacji (scratch nagowywało regresję scratch w benchmarku)") + ap.add_argument("--batch-size", type=int, default=16) + ap.add_argument("--max-iters", type=int, default=4000) + ap.add_argument("--lr", type=float, default=1e-4) + ap.add_argument("--temp", type=float, default=0.85) + ap.add_argument("--topk", type=int, default=18) + ap.add_argument("--new", type=int, default=420) + ap.add_argument("--prompt-chars", type=int, default=96) + ap.add_argument("--seed", type=int, default=42) + ap.add_argument("--split-seed", type=int, default=42) + ap.add_argument("--eval-interval", type=int, default=200) + ap.add_argument("--eval-samples", type=int, default=32) + ap.add_argument("--out", default="data/models/student_rkl_ckpt.pt") + ap.add_argument("--last", default="data/models/student_rkl_last.pt", + help="ostatni checkpoint (niezależnie od val KL) — do porównania z best") + ap.add_argument("--losslog", default="data/models/student_rkl_loss.csv") + a = ap.parse_args() + + sys.stdout.reconfigure(encoding="utf-8") + device = "cuda" if torch.cuda.is_available() else "cpu" + use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported() + ctx = torch.autocast(device_type="cuda", dtype=torch.bfloat16) if use_bf16 else nullcontext() + torch.manual_seed(a.seed) + random.seed(a.seed) + print(f"urządzenie: {device} | bf16: {use_bf16} | seed: {a.seed} | " + f"temp {a.temp} topk {a.topk} new {a.new} | alpha {a.alpha} | {a.direction}") + + # --- student --- + sck = torch.load(a.student, map_location=device, weights_only=False) + stoi, itos, scfg = sck["stoi"], sck["itos"], sck["config"] + student = GPT(scfg); student.load_state_dict(sck["model"]); student.train().to(device) + V = scfg.vocab_size + print(f"student: {a.student} | vocab {V} | block {scfg.block_size} | " + f"parametry {student.num_params():,}") + + # --- nauczyciele: domena -> (model, mapa idx nauczyciel->student) --- + with open(a.domains, encoding="utf-8") as f: + domains_meta = json.load(f) + teachers, t2s_map, meters = {}, {}, {} + for spec in a.teachers: + dom, path = spec.split("=", 1) + tck = torch.load(path, map_location=device, weights_only=False) + if set(tck["stoi"]) != set(stoi): + sys.exit(f"nauczyciel {path}: słownik różni się od studenta — wspólny słownik wymagany " + f"(brak: {sorted(set(stoi) - set(tck['stoi']))}; " + f"dodatkowe: {sorted(set(tck['stoi']) - set(stoi))})") + t2s_map[dom] = torch.tensor([stoi[tck["itos"][i]] for i in range(len(tck["itos"]))], + device=device) + tm = GPT(tck["config"]); tm.load_state_dict(tck["model"]); tm.eval().to(device) + for p in tm.parameters(): + p.requires_grad_(False) + hit = [d for n, d in domains_meta.items() + if isinstance(d, dict) and d.get("type") and dom in d["type"].lower()] + if not hit: + sys.exit(f"domena '{dom}': brak meter w {a.domains} (wpis type zawierający '{dom}')") + meters[dom] = hit[0]["meter"] + teachers[dom] = tm + print(f"teacher {dom}: {path} | meter {meters[dom]} | block {tck['config'].block_size}") + order = list(teachers) + + # --- pule promptów: trening z tr, ewaluacja z val (czysta); scratch kubełkuje po + # długości nagłówka, żeby batch miał równe długości promptów --- + print(f"czytam {a.csv} ...") + rows = load_rows(a.csv) + tr, va = group_split(rows, a.split_seed) + task_list = ["continuation", "scratch"] if a.tasks == "both" else [a.tasks] + pools_tr, pools_va, scratch_buckets = {}, {}, {} + for dom in order: + ok = lambda i: (dom in rows[i]["type"].lower() + and rows[i]["meter"].strip() == meters[dom] + and len(rows[i]["body"]) >= MIN_POOL_BODY + and all(c in stoi for c in build_prompt(rows[i], "continuation", a.prompt_chars))) + pools_tr[dom] = [i for i in tr if ok(i)] + pools_va[dom] = [i for i in va if ok(i)] + buckets = {} + for i in pools_tr[dom]: + buckets.setdefault(len(build_prompt(rows[i], "scratch", a.prompt_chars)), []).append(i) + scratch_buckets[dom] = buckets + print(f" {dom}: pula treningowa {len(pools_tr[dom])}, walidacyjna {len(pools_va[dom])}, " + f"kubełki scratch: {len(buckets)}") + if not pools_tr[dom] or not pools_va[dom]: + sys.exit(f"domena {dom}: pusta pula") + rng = random.Random(f"{a.seed}|eval") + eval_ids = {d: rng.sample(pools_va[d], min(a.eval_samples, len(pools_va[d]))) for d in order} + + def rollout_loss(dom, task, ids): + """(kl, tlp, ent, n) na rolloutach z puli ids. Scratch grupuje po długości nagłówka + (równe długości promptów w batchu); średnie ważone liczbą pozycji.""" + if task == "continuation": + chunks = [(ids, a.prompt_chars)] + else: + by_len = {} + for i in ids: + by_len.setdefault(len(build_prompt(rows[i], "scratch", a.prompt_chars)), []).append(i) + chunks = [(v, L) for L, v in sorted(by_len.items())] + kl_s, tlp_s, ent_s, n_s = 0.0, 0.0, 0.0, 0 + for chunk, plen in chunks: + enc = torch.tensor([[stoi[c] for c in build_prompt(rows[i], task, a.prompt_chars)] + for i in chunk], device=device) + gen = student.generate(enc, a.new, temperature=a.temp, top_k=a.topk) + kl, tlp, ent, n = window_loss(gen, plen, student, teachers[dom], + t2s_map[dom], a.alpha, a.direction, ctx, device) + kl_s += float(kl) * n; tlp_s += float(tlp) * n; ent_s += float(ent) * n; n_s += n + if n_s == 0: + return 0.0, 0.0, 0.0, 1 + return kl_s / n_s, tlp_s / n_s, ent_s / n_s, n_s + + @torch.no_grad() + def evaluate(): + """KL / log-prob nauczyciela / entropia studenta na utwalonych val promptach, + per domena i per zadanie.""" + out = {} + for dom in order: + student.eval() + for task in task_list: + out[(dom, task)] = rollout_loss(dom, task, eval_ids[dom]) + student.train() + return out + + opt = torch.optim.AdamW(student.parameters(), lr=a.lr, betas=(0.9, 0.99), weight_decay=0.1) + + def lr_at(it): + warmup = 100 + if it < warmup: + return a.lr * it / warmup + r = (it - warmup) / max(1, a.max_iters - warmup) + return a.lr * 0.1 + 0.5 * a.lr * 0.9 * (1 + math.cos(math.pi * r)) + + log = [("iter", "train_loss") + + tuple(c for d in order for t in task_list + for c in (f"kl_{d}_{t}", f"tlp_{d}_{t}", f"ent_{d}_{t}"))] + best_kl, t0 = float("inf"), time.time() + for it in range(a.max_iters + 1): + if it % a.eval_interval == 0 or it == a.max_iters: + ev = evaluate() + avg = sum(v[0] for v in ev.values()) / len(ev) + row = [it, ""] + [c for d in order for t in task_list for c in ev[(d, t)]] + log.append(tuple(row)) + line = " | ".join(f"{d}/{t[:4]} kl {ev[(d, t)][0]:.3f} tlp {ev[(d, t)][1]:.3f} " + f"ent {ev[(d, t)][2]:.3f}" for d in order for t in task_list) + print(f"iter {it:4d} | val KL avg {avg:.3f} | {line} | {time.time()-t0:.0f}s") + if avg < best_kl: + best_kl = avg + student.eval() + torch.save({"model": student.state_dict(), "config": scfg, + "stoi": stoi, "itos": itos, "val_loss": best_kl}, a.out) + student.train() + if it == a.max_iters: + student.eval() + torch.save({"model": student.state_dict(), "config": scfg, + "stoi": stoi, "itos": itos, "val_loss": avg}, a.last) + student.train() + break + for g in opt.param_groups: + g["lr"] = lr_at(it) + dom = order[it % len(order)] # rotacja per batch + task = task_list[(it // len(order)) % len(task_list)] # rotacja zadań nad domenami + if task == "continuation": + ids = [random.choice(pools_tr[dom]) for _ in range(a.batch_size)] + plen = a.prompt_chars + else: # scratch: batch z jednego kubełka długości + plen = random.choice(sorted(scratch_buckets[dom])) + ids = [random.choice(scratch_buckets[dom][plen]) for _ in range(a.batch_size)] + enc = torch.tensor([[stoi[c] for c in build_prompt(rows[i], task, a.prompt_chars)] + for i in ids], device=device) + student.eval() # rollouty bez dropoutu (jak w benchmarku) + with torch.no_grad(), ctx: + seq = student.generate(enc, a.new, temperature=a.temp, top_k=a.topk) + student.train() + loss, _, _, _ = window_loss(seq, plen, student, teachers[dom], + t2s_map[dom], a.alpha, a.direction, ctx, device) + opt.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0) + opt.step() + if it % 20 == 0: + print(f" it {it:4d} [{dom:5s}/{task[:4]}] loss {float(loss.detach()):.4f} | {time.time()-t0:.0f}s") + + with open(a.losslog, "w", encoding="utf-8") as f: + f.write("\n".join(",".join(str(x) for x in r) for r in log)) + print(f"\ngotowe. best val KL: {best_kl:.3f}") + print(f"checkpoint best -> {a.out} | last -> {a.last} | krzywa -> {a.losslog}") + + +if __name__ == "__main__": + main()