diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..041877507bf90a95cf84eb77c071627f20887f29 100644 --- a/.gitattributes +++ b/.gitattributes @@ -33,3 +33,51 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00000486_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00000607_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00001171_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00001288_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00001996_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00002166_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00002471_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00002741_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00002795_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00002878_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00003027_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00003738_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00003839_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00003914_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004113_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004291_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004450_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004480_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004656_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004788_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00004822_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00005806_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00005939_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/clean/00006037_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00000486_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00000607_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00001171_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00001288_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00001996_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00002166_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00002471_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00002741_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00002795_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00002878_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00003027_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00003738_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00003839_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00003914_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004113_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004291_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004450_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004480_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004656_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004788_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00004822_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00005806_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00005939_a0.wav filter=lfs diff=lfs merge=lfs -text +demo_samples/val/degraded/00006037_a0.wav filter=lfs diff=lfs merge=lfs -text diff --git a/README.md b/README.md index fdc8c3880523572f4bc187ea8b1eafaf5662a9bf..14156f342c2d70669b44fb44aec0f53a2fe5a202 100644 --- a/README.md +++ b/README.md @@ -1,13 +1,23 @@ --- -title: Stem Restoration -emoji: ๐Ÿ  -colorFrom: blue -colorTo: blue +title: Stem Restoration + Generation +emoji: ๐ŸŽ›๏ธ +colorFrom: gray +colorTo: yellow sdk: gradio sdk_version: 6.19.0 -python_version: '3.13' +python_version: "3.10" app_file: app.py pinned: false --- -Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference +# Stem Restoration + Generation + +Upload a muffled or lo-fi song and get a cleaned-up version back. Under the hood we split it into instruments, then: + +- **Restorer** โ€” *cleans up* each muffled/damaged instrument so it sounds clear again (it fixes what's there; it doesn't add new parts). +- **Generator** โ€” *invents* a missing instrument (mainly bass) that fits the song, like an AI session musician. +- **GAN 'judge'** โ€” rates how *real* the final mix sounds; the restorer and generator compete to fool it, pushing the result toward genuinely real audio instead of a dull average. + +Everything runs in the compact **SAME-L** neural-audio latent space, on CPU. + +Active models โ€” restorers: ['restorer_attn_advramp', 'restorer_attn_w003_gan'] ยท generators: ['gen_advramp_v1', 'gen_distvar_baseline'] diff --git a/app.py b/app.py new file mode 100644 index 0000000000000000000000000000000000000000..5e06417a6549c5a77c222efdd5602efbfc6072b9 --- /dev/null +++ b/app.py @@ -0,0 +1,8 @@ +# HF Space entry. Paths resolve repo-relative (see restoflow/app.py). +import os +os.environ.setdefault('RESTOFLOW_DEVICE', 'cpu') +from restoflow.app import build_ui, models +import threading +demo = build_ui() +threading.Thread(target=models, daemon=True).start() +demo.queue().launch(server_name='0.0.0.0', server_port=7860, ssr_mode=False) diff --git a/demo_samples/val/clean/00000486_a0.wav b/demo_samples/val/clean/00000486_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..40da96ffd9892eaf58ee7e515eb5c9a11bcc1213 --- /dev/null +++ b/demo_samples/val/clean/00000486_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba7c7eb749901e59bc59acfc4f29d41ab61195af01449b60230a06e580fd3dc8 +size 529244 diff --git a/demo_samples/val/clean/00000607_a0.wav b/demo_samples/val/clean/00000607_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..3fdb9581567c202560fbbd7225ab7b7efbe6b675 --- /dev/null +++ b/demo_samples/val/clean/00000607_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6135a3bbdeacc42554203a58dac07c5c9d2f6b6a7558e1f5d5b1f4298cd86bd7 +size 529244 diff --git a/demo_samples/val/clean/00001171_a0.wav b/demo_samples/val/clean/00001171_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..2cac1305914809c9cedc191c5e57b7d3b4bbefef --- /dev/null +++ b/demo_samples/val/clean/00001171_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0c671f39e0fbb0acb9a6293ca42ed3dfa3a032cf1ee86a6b451ec134edfa3e49 +size 529244 diff --git a/demo_samples/val/clean/00001288_a0.wav b/demo_samples/val/clean/00001288_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..9c94fa8522c2a3b4f4767131f447128fef5ec040 --- /dev/null +++ b/demo_samples/val/clean/00001288_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dc1b550cd01ff246bad6361a83f3453831d7ffc0d2d368a97cd6119dab27cd1e +size 529244 diff --git a/demo_samples/val/clean/00001996_a0.wav b/demo_samples/val/clean/00001996_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..21b13ec3e7919bab90ca5d82b1b044c055e0bec2 --- /dev/null +++ b/demo_samples/val/clean/00001996_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:585fd5c504f6795abd94bb3445528a280cb75431859d35e128b8a8014c85034f +size 529244 diff --git a/demo_samples/val/clean/00002166_a0.wav b/demo_samples/val/clean/00002166_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..2c978333c3d2c432167afc940501e996027f5389 --- /dev/null +++ b/demo_samples/val/clean/00002166_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d4ab35e73315d7def12eb2224f6f6426092b0339b2af797dd5e49f5a39abb18e +size 529244 diff --git a/demo_samples/val/clean/00002471_a0.wav b/demo_samples/val/clean/00002471_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..7eb8bb9d58a126551917b5bdcbea51463742471c --- /dev/null +++ b/demo_samples/val/clean/00002471_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b5db4ce9487261baedb2bf2c9e8b39c48db586275028e93d5f3b0ad3eeacdda +size 529244 diff --git a/demo_samples/val/clean/00002741_a0.wav b/demo_samples/val/clean/00002741_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..4f7089648859c82544687e49dadaa1a5970ace64 --- /dev/null +++ b/demo_samples/val/clean/00002741_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f493cddf753a26c67508f371d3d7dbfe601af74cfbb512e417a9b0a2b6703782 +size 529244 diff --git a/demo_samples/val/clean/00002795_a0.wav b/demo_samples/val/clean/00002795_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..ce47c3e8f9df17b5db1aa62ab9c31d51f6c13561 --- /dev/null +++ b/demo_samples/val/clean/00002795_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a2b57c8dcec7d33510168a6bb484646653413c1917fc3066d5a3b1e83647b74c +size 529244 diff --git a/demo_samples/val/clean/00002878_a0.wav b/demo_samples/val/clean/00002878_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..794751d0f8a474b0b15471cf071958728bfeb769 --- /dev/null +++ b/demo_samples/val/clean/00002878_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5324212d2e7562cba0554ca783ac0202d1c615f5e70432680033b2c6f8e1cfa4 +size 529244 diff --git a/demo_samples/val/clean/00003027_a0.wav b/demo_samples/val/clean/00003027_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..f1f1cd2c3dfc1bba2bbd901c9758ccddd31dbe43 --- /dev/null +++ b/demo_samples/val/clean/00003027_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:85a2a33f7987ab5fa94ca9167002a22b2375d06f899c766744ada9b3c2dd5fa6 +size 529244 diff --git a/demo_samples/val/clean/00003738_a0.wav b/demo_samples/val/clean/00003738_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..20a7b96bbe7a1671d8ec7059f569914902669513 --- /dev/null +++ b/demo_samples/val/clean/00003738_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47f2dc38ee7ff1aab08d0acec4871dcd326049807f5ebc650031e4b9cb8960f3 +size 529244 diff --git a/demo_samples/val/clean/00003839_a0.wav b/demo_samples/val/clean/00003839_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..75f57c6e87d3be8d756f1b463a8c0f42f81a4797 --- /dev/null +++ b/demo_samples/val/clean/00003839_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8c7b521f783aac7ab368a974988d8d47db42e42b5c8d28e0d847b039a4971f92 +size 529244 diff --git a/demo_samples/val/clean/00003914_a0.wav b/demo_samples/val/clean/00003914_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..ec299745570c4d155d27a7be109c4e5764934e59 --- /dev/null +++ b/demo_samples/val/clean/00003914_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f40ea9c5e7ffffed0e86fb1b4b1f2b777004bc1300fbb042d55dc7154bac6f3d +size 529244 diff --git a/demo_samples/val/clean/00004113_a0.wav b/demo_samples/val/clean/00004113_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..f031e467ecd9dcc6a2d9a6c71c660a7ea9713b75 --- /dev/null +++ b/demo_samples/val/clean/00004113_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:59552c0b6f5a4d8c77b7f96f98d404098849189a68edd0379eff4be913854651 +size 529244 diff --git a/demo_samples/val/clean/00004291_a0.wav b/demo_samples/val/clean/00004291_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..8b132740f473cca623eece1e9da012b1deac9a66 --- /dev/null +++ b/demo_samples/val/clean/00004291_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:387465d96ad33e75ee3574bfb55615e5babb20d608d14a6c92530b3eb3844e3b +size 529244 diff --git a/demo_samples/val/clean/00004450_a0.wav b/demo_samples/val/clean/00004450_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..3962b9f9a26babba38de41cf9597199bc639f55e --- /dev/null +++ b/demo_samples/val/clean/00004450_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:abd469d5ce2c0cbf73fbe48d94439e15289e68c6d7ca7ccd1f4582ccf97e008e +size 529244 diff --git a/demo_samples/val/clean/00004480_a0.wav b/demo_samples/val/clean/00004480_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..8a0ae0b27d1534c4dae40e923e9c471f52d25659 --- /dev/null +++ b/demo_samples/val/clean/00004480_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:801346ea55485b41c43719796529eba538b0df3d3b9273f29f9a063125e942f2 +size 529244 diff --git a/demo_samples/val/clean/00004656_a0.wav b/demo_samples/val/clean/00004656_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..72f000526e0625b6eee288f8062104da3c4f38f8 --- /dev/null +++ b/demo_samples/val/clean/00004656_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:12158bd5a77be1a9e2a6f8afd48a683f70755c101e418845b497b66450c45559 +size 529244 diff --git a/demo_samples/val/clean/00004788_a0.wav b/demo_samples/val/clean/00004788_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..e25dc45d2e818f2087696d56484a967506b40cdb --- /dev/null +++ b/demo_samples/val/clean/00004788_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f68def37024b5263acb66e26a5e40eb925c12e4c7898e28b626f01eb4b01b208 +size 529244 diff --git a/demo_samples/val/clean/00004822_a0.wav b/demo_samples/val/clean/00004822_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..75ddf9cd3c6857fbe2d2b2aa119b3dc85cb5a294 --- /dev/null +++ b/demo_samples/val/clean/00004822_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a15b4ec28615d69a4f9a022dc876354783f321d1cab14c216ebe945322a93720 +size 529244 diff --git a/demo_samples/val/clean/00005806_a0.wav b/demo_samples/val/clean/00005806_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..453ccb5d34382780f96d1bb110b20c39f96b4edf --- /dev/null +++ b/demo_samples/val/clean/00005806_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ae98f54ef448f4e1d5748118099e1274e060243a2d84cb39b1eba2651f02ee8d +size 529244 diff --git a/demo_samples/val/clean/00005939_a0.wav b/demo_samples/val/clean/00005939_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..a29184d408ab5fa2feac71d5372d809a8470db47 --- /dev/null +++ b/demo_samples/val/clean/00005939_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:22c25cbb2c8e899b3f5a748dbc146725db36570b90f66f7e7b8ef5c372d8bd05 +size 529244 diff --git a/demo_samples/val/clean/00006037_a0.wav b/demo_samples/val/clean/00006037_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..7cb00dbc63a8e9d9501183de8d5fd9eb051d0b0e --- /dev/null +++ b/demo_samples/val/clean/00006037_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0723683fe3b6df79cba1b5e17b411f1524b54b6e67e1e72ccaad5e12ebaacb04 +size 529244 diff --git a/demo_samples/val/degraded/00000486_a0.wav b/demo_samples/val/degraded/00000486_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..7cc098252124584a9ba215e5e06ef2d61cfc32e1 --- /dev/null +++ b/demo_samples/val/degraded/00000486_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2c05b12beebe4cb8d5ee7862bc7ac2978713be90e8f2ad1d6d02fb63571e369c +size 529244 diff --git a/demo_samples/val/degraded/00000607_a0.wav b/demo_samples/val/degraded/00000607_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..3cad0166816cdacad38eba53bc08abf32cdf780a --- /dev/null +++ b/demo_samples/val/degraded/00000607_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8eaf0efa7f97a28e4e9d9abcfc321ec8dcfc08eac1cbe16e5d04037d3b99f08d +size 529244 diff --git a/demo_samples/val/degraded/00001171_a0.wav b/demo_samples/val/degraded/00001171_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..070db72c0d676f7a6c6d855413d8637d490a8844 --- /dev/null +++ b/demo_samples/val/degraded/00001171_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be91484a919d3bcdac1b0c473b36617674ebadd3ca413f14252477f9dcd66bad +size 529244 diff --git a/demo_samples/val/degraded/00001288_a0.wav b/demo_samples/val/degraded/00001288_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..590095e809a8143524f28b653fc9fdcd0eb86b08 --- /dev/null +++ b/demo_samples/val/degraded/00001288_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9c3d7745bf3a6faebe429794843b2d4cc34c023d363a7a6e4ebf1f7e3b4f22f1 +size 529244 diff --git a/demo_samples/val/degraded/00001996_a0.wav b/demo_samples/val/degraded/00001996_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..a3664cac02287f0db48bbde205ca7d1da5ff1dfd --- /dev/null +++ b/demo_samples/val/degraded/00001996_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a563144c7f47b844d17b680911e7bafff19d9502bb9077a14464259553539ce0 +size 529244 diff --git a/demo_samples/val/degraded/00002166_a0.wav b/demo_samples/val/degraded/00002166_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..1bcec0d2d88cb7545308af018bc10c9e28c425c0 --- /dev/null +++ b/demo_samples/val/degraded/00002166_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:49f630640c5b454d19a8868ca27b75fa521e100792806d7773a1783e7cb1d80b +size 529244 diff --git a/demo_samples/val/degraded/00002471_a0.wav b/demo_samples/val/degraded/00002471_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..f1af1335bd14a873df04477862459d51dbbaa411 --- /dev/null +++ b/demo_samples/val/degraded/00002471_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ed61afd278300d5f80d05da146f5e7445c8b190883fe00b56d899e188a83454d +size 529244 diff --git a/demo_samples/val/degraded/00002741_a0.wav b/demo_samples/val/degraded/00002741_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..6f1d5a782480da7d2289ea1369d7ee137fc22a72 --- /dev/null +++ b/demo_samples/val/degraded/00002741_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c4b03514fbfbe688d6589289ddf855b13a03a09c7d8136caca2d752b9d1618ef +size 529244 diff --git a/demo_samples/val/degraded/00002795_a0.wav b/demo_samples/val/degraded/00002795_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..7792590e4b1adfa2558517b47280949cce107e53 --- /dev/null +++ b/demo_samples/val/degraded/00002795_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:89aa2a64253cb085cd77f0f312ff0397bf1616b0bc2770ac0bf2f1e546f82652 +size 529244 diff --git a/demo_samples/val/degraded/00002878_a0.wav b/demo_samples/val/degraded/00002878_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..49891d12d2836af688e45b6d29f90e90bc24cc0e --- /dev/null +++ b/demo_samples/val/degraded/00002878_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e3e4ac53bb2f7600c9187ff44768d6ad1ff930d62993bd95a1a0403594ca25b +size 529244 diff --git a/demo_samples/val/degraded/00003027_a0.wav b/demo_samples/val/degraded/00003027_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..79223feae30229fcf61820a5edbd65ababed68e7 --- /dev/null +++ b/demo_samples/val/degraded/00003027_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:adc54f66e871df3e62fa4d42e1f1ed3ecd4df33f4713c600eb4d2b9310bb4880 +size 529244 diff --git a/demo_samples/val/degraded/00003738_a0.wav b/demo_samples/val/degraded/00003738_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..4675cbd08fc2c4a7850512d59a483ab47d1360ce --- /dev/null +++ b/demo_samples/val/degraded/00003738_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:29806184c53a3de7212dfbb76df7f5ef14f008cbb19e1162c033da53e0ceb5ec +size 529244 diff --git a/demo_samples/val/degraded/00003839_a0.wav b/demo_samples/val/degraded/00003839_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..a6d87752ea1f9319e1337d92e63d2272a2778cb7 --- /dev/null +++ b/demo_samples/val/degraded/00003839_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba55496b401498c3fc272f8c9eadc253917c8e2fa514ca0ad409e246cc4c5318 +size 529244 diff --git a/demo_samples/val/degraded/00003914_a0.wav b/demo_samples/val/degraded/00003914_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..e1e27ae53b51dff65f8d2dc47f4feb62329a50d5 --- /dev/null +++ b/demo_samples/val/degraded/00003914_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7ff5cfcb1340481a53abf21e09b1f4278f723b67cc1f275e00b45afb1a83f8d9 +size 529244 diff --git a/demo_samples/val/degraded/00004113_a0.wav b/demo_samples/val/degraded/00004113_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..5669acbc94cd55072d7ad3f7c54ccb31065dbf63 --- /dev/null +++ b/demo_samples/val/degraded/00004113_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ff6748c6e5c75656df78dd890e673e6a917293f676accdf758281458988a5d4e +size 529244 diff --git a/demo_samples/val/degraded/00004291_a0.wav b/demo_samples/val/degraded/00004291_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..126e86e9b2c0c7b1d293bd0acb2295dd3412354e --- /dev/null +++ b/demo_samples/val/degraded/00004291_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09dfa918d16a8e3e72d04b9e953d02506f4703cafb23842749701c7a0380ef34 +size 529244 diff --git a/demo_samples/val/degraded/00004450_a0.wav b/demo_samples/val/degraded/00004450_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..a59697d5d9cddcc4ae923eb15b42f9128122f32e --- /dev/null +++ b/demo_samples/val/degraded/00004450_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:40273b7a60056a6aa60e4f729c98a846d1bb711fe0f38340cf0ae2fd84ef1a2b +size 529244 diff --git a/demo_samples/val/degraded/00004480_a0.wav b/demo_samples/val/degraded/00004480_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..50dc7e3af24f77079e2e30818bebe1fd01b005cf --- /dev/null +++ b/demo_samples/val/degraded/00004480_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7508c3b40da298c02afebfcf0bf795b8b02880899a21a937601b79e34968199a +size 529244 diff --git a/demo_samples/val/degraded/00004656_a0.wav b/demo_samples/val/degraded/00004656_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..2c57d2074cdb300e79a89a6a0b9304dd8fecc2a9 --- /dev/null +++ b/demo_samples/val/degraded/00004656_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fd74f5241cbd9b308a524de61e89235e43b7c061fab272d3d97c54adb97408fc +size 529244 diff --git a/demo_samples/val/degraded/00004788_a0.wav b/demo_samples/val/degraded/00004788_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..34e2090314291be69796c124353a404ff9292ad4 --- /dev/null +++ b/demo_samples/val/degraded/00004788_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a059f74c06d98e128d98d4f1c3cb293b512797c1cca2a5c2ec4b1c3e0920cd4f +size 529244 diff --git a/demo_samples/val/degraded/00004822_a0.wav b/demo_samples/val/degraded/00004822_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..e19b5b308e2e42308a61db560f5f986c7e504682 --- /dev/null +++ b/demo_samples/val/degraded/00004822_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:88fd2ea8bdc8ef99d1dc7e0fe2d0ffc16dec572fefc6b324840d938278c83712 +size 529244 diff --git a/demo_samples/val/degraded/00005806_a0.wav b/demo_samples/val/degraded/00005806_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..fa402fa45699be73fd55c4c3e03b5bb096f30457 --- /dev/null +++ b/demo_samples/val/degraded/00005806_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3207715856781974aa89118a36f2866a250315d9d5f18d2aa09ac7725369b725 +size 529244 diff --git a/demo_samples/val/degraded/00005939_a0.wav b/demo_samples/val/degraded/00005939_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..dfd8c6b4f1ce14de93fc8a8e3aec18706b085a66 --- /dev/null +++ b/demo_samples/val/degraded/00005939_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54d4c9ae38a007cca22ce3ec496203b0d52ec65464c79772341cdec40a8e2d8d +size 529244 diff --git a/demo_samples/val/degraded/00006037_a0.wav b/demo_samples/val/degraded/00006037_a0.wav new file mode 100644 index 0000000000000000000000000000000000000000..76375ae1c0a538de3cd40401ab94c2f7563f1694 --- /dev/null +++ b/demo_samples/val/degraded/00006037_a0.wav @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6dd04f75426e54c6b871223aca49be06c4a250223f44a5fc829ef005d90598fc +size 529244 diff --git a/demo_samples/val_cache/00000486_a0/latents/clean/other.pt b/demo_samples/val_cache/00000486_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..bf4450cdbd89639d8b32d9af274a0515cfaab656 --- /dev/null +++ b/demo_samples/val_cache/00000486_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9cb1ab35c64e63a66156fcaad988ff893081d037492f0090c92a19a21d7ae93b +size 18395 diff --git a/demo_samples/val_cache/00000486_a0/latents/degraded/other.pt b/demo_samples/val_cache/00000486_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..2c9d8941bf14caaba8276436893e44faadbb7949 --- /dev/null +++ b/demo_samples/val_cache/00000486_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4213db750d02fd8d74ae11a45a0ad367e132fcd9f78a9e0a31850e6e48ba44c2 +size 18395 diff --git a/demo_samples/val_cache/00000486_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00000486_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..348fa2bbb5c0a3fcd4ab5b5bb274a62c60eedfd7 --- /dev/null +++ b/demo_samples/val_cache/00000486_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fe4aeca6a26f1147a11cbff728fcfb4e5e6068d0ea39b6223e3c3bda1ab6ae62 +size 18508 diff --git a/demo_samples/val_cache/00000486_a0/latents/restored/other.pt b/demo_samples/val_cache/00000486_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..31bb96304cb313d6f2a5883b7cc691b7f6d04111 --- /dev/null +++ b/demo_samples/val_cache/00000486_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ec6fa5b6ab2ab6fa2a1cf7a5b3c09114685aa24306a74f31af48b16afaf0ca8e +size 18395 diff --git a/demo_samples/val_cache/00000607_a0/latents/clean/bass.pt b/demo_samples/val_cache/00000607_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..41b022f4383b61595e5bf62ac6710f683afc85b7 --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5a30f779cdaff0247eb90bb36fe1c5c2a8ed4bd1a2c47d8672e7b41a4602921a +size 18388 diff --git a/demo_samples/val_cache/00000607_a0/latents/clean/drums.pt b/demo_samples/val_cache/00000607_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..15ba7eae6bf2252706d09b0cb9d2e6566e0e9c7b --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72c6c882dd1db8b114427f9afbf1ef6341861b5d06936fcb5610f78df21809c0 +size 18395 diff --git a/demo_samples/val_cache/00000607_a0/latents/clean/other.pt b/demo_samples/val_cache/00000607_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..69ce7c050fe824b9a8014370e1245d53be13fd2e --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5f4eb37ff2e36a1cfe2e0b24abdcef6b3a678638eebb833418f12f682b15c1c1 +size 18395 diff --git a/demo_samples/val_cache/00000607_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00000607_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..64309f55bff8111bc4e60717c39d9fcc4948a4f6 --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0807008d3ef015eecddbcf891215fed6ab084756007e193543e0a920537074f4 +size 18402 diff --git a/demo_samples/val_cache/00000607_a0/latents/degraded/other.pt b/demo_samples/val_cache/00000607_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..f6ef2a57e7c8b8687fac29e873fecc6cb145e8dc --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:11b75e4d0d0da466b3e3967be493563f9a2b445e7f526e2b6dd8092447bd24f3 +size 18395 diff --git a/demo_samples/val_cache/00000607_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00000607_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..a6180365d9945f7deb9170a6bb6ba3356160c3f6 --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23f027cdf86e5e0306fe39e8712fafd21d0205bc58d0f04d4bb0b42a3b249e9a +size 18402 diff --git a/demo_samples/val_cache/00000607_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00000607_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..a455147c601cb3fee295cd1ef455a86ac239f8ef --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f6e89efa3f8b8cc22896685cc095a5b9dba6a4c9527e301112cdc87f3d96f2dd +size 18508 diff --git a/demo_samples/val_cache/00000607_a0/latents/restored/other.pt b/demo_samples/val_cache/00000607_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..4265e32fcb686171261bc93699b564dec5a2c76e --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9e5b3f6de4b4c4ffa83e26b5fdcae72f18933ab4c7b842019ae9e6a5697e726a +size 18395 diff --git a/demo_samples/val_cache/00000607_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00000607_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..acfbec110ef13f954d67c34622976ce7aa94ea5c --- /dev/null +++ b/demo_samples/val_cache/00000607_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:da37a0ee16a7d7b98a551db31a8c96519344f5e9a2e931f849c0e90dd9a1712a +size 18402 diff --git a/demo_samples/val_cache/00001171_a0/latents/clean/bass.pt b/demo_samples/val_cache/00001171_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..c3008f83bb672ae49b2e84666712b17bed52fbbe --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fad11aa4d83fd39c509eed070b971f54614b3590beb07683086f28c893065f93 +size 18388 diff --git a/demo_samples/val_cache/00001171_a0/latents/clean/drums.pt b/demo_samples/val_cache/00001171_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..42de4b33a88fe78e3a4c06cef022b6f02039196f --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ce55f0dca839f068e225859e7aee9e5f2daff1b3e3762faaa16eb422ee50cbd8 +size 18395 diff --git a/demo_samples/val_cache/00001171_a0/latents/clean/other.pt b/demo_samples/val_cache/00001171_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..c83e3cbcc92288bd4a68035c37bbe1b1d5196921 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:30c84d231ee16125ad625632748285ae7fb3d2719e018d82dbedd0a1161d2252 +size 18395 diff --git a/demo_samples/val_cache/00001171_a0/latents/degraded/bass.pt b/demo_samples/val_cache/00001171_a0/latents/degraded/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..b9927c85adac79ab1de06f32c2e2523d20c4d4a1 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/degraded/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:778708025de0ed2739bc761da066f27201d75a10c4b6e0fcb9f052d84eb6a558 +size 18388 diff --git a/demo_samples/val_cache/00001171_a0/latents/degraded/drums.pt b/demo_samples/val_cache/00001171_a0/latents/degraded/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..f005aa11e2baf2f82986f8b6f94f7be177ed8993 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/degraded/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa2442ca54028987f093fd519171e680b1810b160c18ba1aab8ca2e980e23eb0 +size 18395 diff --git a/demo_samples/val_cache/00001171_a0/latents/degraded/other.pt b/demo_samples/val_cache/00001171_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..5a243f941b4ade3c2c68d338f034717dc598832d --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8b99e06a24e3c907b5a6213160e0d70840e10a2abd745b85355adb991ce9aaf9 +size 18395 diff --git a/demo_samples/val_cache/00001171_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00001171_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..9805f1e5db5247253d6c3cab7be29eeead5278b5 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:214f6b28a6c8d03b7abc98cbd67356fe27972cc137f0b839624f8c667288ab28 +size 18508 diff --git a/demo_samples/val_cache/00001171_a0/latents/restored/bass.pt b/demo_samples/val_cache/00001171_a0/latents/restored/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..dc915de5f64dd1ee424ef1fdfc590798bcb37372 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/restored/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d02bd366436193d5d9b51f8572a46923c26fffa450ad83b14795401d8371a7d8 +size 18388 diff --git a/demo_samples/val_cache/00001171_a0/latents/restored/drums.pt b/demo_samples/val_cache/00001171_a0/latents/restored/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..f57050a68111d6010b925d745454f645d904d3bd --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/restored/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:10d4c722776fed4f34a596f00e3db17dce7f36b57667f5aac62ef7e847a01383 +size 18395 diff --git a/demo_samples/val_cache/00001171_a0/latents/restored/other.pt b/demo_samples/val_cache/00001171_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..f83a08e70610a8a3852c775abf52b961a0c99d73 --- /dev/null +++ b/demo_samples/val_cache/00001171_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:005385ebed0017576209a2f52ee08400ad96e14ded16bb37b78e84d235cdc3f0 +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/clean/drums.pt b/demo_samples/val_cache/00001288_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..b1b7ad4807ea3b366dea4526a069b39dfb7d6fd9 --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb309f2a2f9b6540b285b452c7244bf6211a784078ae06a2ea3c38102a5dd292 +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/clean/other.pt b/demo_samples/val_cache/00001288_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..99513609ac85b0a504ad8a0eca773679d0b50dc1 --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ba5b541b317bd3c6dcc109ac5e294831571b310e0bd24e0d2de6d1384cce078f +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/degraded/drums.pt b/demo_samples/val_cache/00001288_a0/latents/degraded/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..bbb6ac8ca4ce8343becd73a55b9d75f6a50423c9 --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/degraded/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d71dc5b797e8b60c5655aca527dd87827d430a187448874f61503d92a6f637c +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/degraded/other.pt b/demo_samples/val_cache/00001288_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..cd101e02c8f0dc27c40537b2e9cf4f779ac3fb75 --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a1faf34915cfdafad3e868ae3c6dc9f68e4e85f92f88d4c42ac0e497f212bd45 +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00001288_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..86ba682288474d5e19b8da1f4ef70c5fbce3f9fb --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:508932bc66593562a47ddaa530451cf1a2f0c61d4bf1f7f6cd0f815bd1e2d118 +size 18508 diff --git a/demo_samples/val_cache/00001288_a0/latents/restored/drums.pt b/demo_samples/val_cache/00001288_a0/latents/restored/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..239800c59fd557e3cb2e70c6657fc411bc8d6461 --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/restored/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72614752d194009a1c8847f7598a3a449eb2ab0de5afed24d34ddc57d20f9338 +size 18395 diff --git a/demo_samples/val_cache/00001288_a0/latents/restored/other.pt b/demo_samples/val_cache/00001288_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..7a72fef1bac6db7f2b2ef43b7cf3c5ba0d1974ff --- /dev/null +++ b/demo_samples/val_cache/00001288_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8536d38112cf0ee32a5c73ea98e61f10c33c019f937c8c4b54b95fd6f288a22c +size 18395 diff --git a/demo_samples/val_cache/00001996_a0/latents/clean/other.pt b/demo_samples/val_cache/00001996_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..dfd0046b5f82ca5242fc828413e69e30edddef4a --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2df236e4061c30d0818d794d2c37b7e96b225d8410b3dafa189a6f4061ed8398 +size 18395 diff --git a/demo_samples/val_cache/00001996_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00001996_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..65f741b64614b038f911c0d1a69fc6ab6ef9cc50 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ef77c3d3c3089c535a72830d1b8591c84576321de609fba9d1f257b388f12f56 +size 18402 diff --git a/demo_samples/val_cache/00001996_a0/latents/degraded/other.pt b/demo_samples/val_cache/00001996_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..27393376dd580a6b8204491dc997e2f689cae731 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97fc2d6f92f2382643dd54a95fde6335ecfcb7585f832daab842777281ef5aaa +size 18395 diff --git a/demo_samples/val_cache/00001996_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00001996_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..e87bdb4ae62839563ef0bda5999552977b66ae75 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:43ef664bac3d43d04fcb6f7a0aee92532ed3a1afa62b8febaa6300246310adb0 +size 18402 diff --git a/demo_samples/val_cache/00001996_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00001996_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..5f9114b818ca6a844cbc3594d6c4eaf5403e7644 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf58140dc7e4c7165299306d7f91acd7f5cdfe9ae3326e4d2d9876f54207f529 +size 18508 diff --git a/demo_samples/val_cache/00001996_a0/latents/restored/other.pt b/demo_samples/val_cache/00001996_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..0b5e5e01ebaf4bd74e246ff61059bb6aae08e0a3 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aefb59f5cdd1b6599afa39114d5731949368110edbe3dd7d42c33412d28ed72f +size 18395 diff --git a/demo_samples/val_cache/00001996_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00001996_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..bfbcd48aa1510e8a86d41df74d381865c7c2a533 --- /dev/null +++ b/demo_samples/val_cache/00001996_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f5e3ffc26f9dcde1b26fabf432cf9d54a0ed66bfef3221c12ae9df0c84157ecd +size 18402 diff --git a/demo_samples/val_cache/00002166_a0/latents/clean/bass.pt b/demo_samples/val_cache/00002166_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..7fbfcbd38ed5eeff13107ad8b9b353dc43a4152e --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0863c29457ac6f517107779754cc3dbad12d398c3de0de7941a8a567e8ba9c62 +size 18388 diff --git a/demo_samples/val_cache/00002166_a0/latents/clean/other.pt b/demo_samples/val_cache/00002166_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..12251b0911cd41491a3a89e7a3b7a9c34e0150d4 --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bbadb49e5aa92cdcbf19191b85faeb0057eeb5357619c6ebaf44c78dede43c42 +size 18395 diff --git a/demo_samples/val_cache/00002166_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00002166_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..bac13319ce5747d598a78999223dbc968cb7153d --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fa20e359a2c6c5abd861cb1f75e62e65c3142e5e6a7ab31ec178bdfd896cb459 +size 18402 diff --git a/demo_samples/val_cache/00002166_a0/latents/degraded/other.pt b/demo_samples/val_cache/00002166_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..0fdd814f8cb35e0ec7f705f3272e60552f34c30b --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8657191288a9abe411f7546156415fbf32fcb85cdf94d71e869690515afab14a +size 18395 diff --git a/demo_samples/val_cache/00002166_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00002166_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..d1f7b2d8c2f992ef52880f43c2030b99eb2e8276 --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1e2a03621190e5e845038ab8a6c4e4f791336a46d78e282a9bc8e0516b906be4 +size 18402 diff --git a/demo_samples/val_cache/00002166_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00002166_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..b9f24c7b74155c27735645f1eed26f982e9a9872 --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3bb4c2ea9ab1ddac4b91b54c6a548ced831b56e2de11ab0a8b2142c07ec5255 +size 18508 diff --git a/demo_samples/val_cache/00002166_a0/latents/restored/other.pt b/demo_samples/val_cache/00002166_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..4a34b930b527caba8f982b69fa4287fdc01ee1cb --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3403b30984836f4de907ac540e526b2add8b123a5a8ad11f9dd7248bc1a73e0 +size 18395 diff --git a/demo_samples/val_cache/00002166_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00002166_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..6ef72dfa2333093ada466fd9bdbb3dea2190a84e --- /dev/null +++ b/demo_samples/val_cache/00002166_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5433fc4a6f7592810a8c14d1ec9db4cd268fd28727e74413f393fc2694591909 +size 18402 diff --git a/demo_samples/val_cache/00002471_a0/latents/clean/drums.pt b/demo_samples/val_cache/00002471_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..bbe0ebd6e0966917d1c76724cc345d172092debc --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:940afaf598b19c4f97a83d952df2deea21bd13cecf68b392af79c1157f4232c6 +size 18395 diff --git a/demo_samples/val_cache/00002471_a0/latents/clean/other.pt b/demo_samples/val_cache/00002471_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..8679af1227ccf08e214dd11bd6eb7b1493cdf712 --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ca6376d89c63cd9780ec19ca02150e094790f4808dc729eb1c52318331e18a1f +size 18395 diff --git a/demo_samples/val_cache/00002471_a0/latents/degraded/drums.pt b/demo_samples/val_cache/00002471_a0/latents/degraded/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..288b241c9e4ec230d75cf939525deadfe9edfdd4 --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/degraded/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ebf517bd24f379b9395a110e52cb86398eb2f2ad676d09796e549bd28c9eaf00 +size 18395 diff --git a/demo_samples/val_cache/00002471_a0/latents/degraded/other.pt b/demo_samples/val_cache/00002471_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..9d6a9cfd73fff9b2e73d7d8d57a6d06d56c5e6ca --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8947c331bd39305df8cb5740a67cc7414a8c7863ae049fd973026d3be2ec8e73 +size 18395 diff --git a/demo_samples/val_cache/00002471_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00002471_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..cccfe10ec5f92d47f107ade2ee9e557b28d9f385 --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8c4911784cbbcab4e740b737740929e72a47724ea8bf09d8b4b879731dd8d64 +size 18508 diff --git a/demo_samples/val_cache/00002471_a0/latents/restored/drums.pt b/demo_samples/val_cache/00002471_a0/latents/restored/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..d27b0c0ceb1cfa82ee6d82db2a52ad20bcdce640 --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/restored/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f1d45f57397b45046bbb27fbc1e14105cda74bce48086af546b55f76c671948 +size 18395 diff --git a/demo_samples/val_cache/00002471_a0/latents/restored/other.pt b/demo_samples/val_cache/00002471_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..0dee8f2a3ce8ee729295558b1cce03a3865b6b1f --- /dev/null +++ b/demo_samples/val_cache/00002471_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:36e337b2b2fffee1e5e9e6a5c0a8ef766530a11d1831427716ae9e0fd2829a4e +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/clean/drums.pt b/demo_samples/val_cache/00002741_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..ee2005689f372ef499c16a0cf1821aa79817c32f --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3ec16a8ea987a5486d59037b8e7dcbd07c88b3a6099d80822ede8458b8dd9475 +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/clean/other.pt b/demo_samples/val_cache/00002741_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..b76e931722488dc2c85f4d6bf9cd6a0b5a40d752 --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:533c6543143adea2deeb3008c141e7f8af3fcd11b6583c659bad99cdba3f51db +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00002741_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..988bf4c69bc6641710358647d0845d1f5698693b --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd5f35358b9805382ee22f12aa8d0eaae91ec5afcc003d9d2a210d5c3a58184c +size 18402 diff --git a/demo_samples/val_cache/00002741_a0/latents/degraded/drums.pt b/demo_samples/val_cache/00002741_a0/latents/degraded/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..5318dd014d5805e0cb6715ae6505b160f470d120 --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/degraded/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d96087b9dc2a71e4df05431820460a60d1084ca55774714ecf7daccdc76acc17 +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/degraded/other.pt b/demo_samples/val_cache/00002741_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..1d767efec011ea0adcaf5c89801867d56bd6fe0a --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8a050c610f5ca7759554824f6806817b756cded8fba6f8f3b30669d67ef3f72c +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00002741_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..56500235684d1c1d80435b43137300c165c681ce --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f660bba73a3bb6771edd1ff01b7a29eff4018a2afad375e243f251ffdb8cc6e5 +size 18402 diff --git a/demo_samples/val_cache/00002741_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00002741_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..f40270a673fc8072d714c78147c11603d8a6d473 --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77f4eab414568a9b0bfc3c284cdded58cee90734a7c3bdf06c65afe62830c944 +size 18508 diff --git a/demo_samples/val_cache/00002741_a0/latents/restored/drums.pt b/demo_samples/val_cache/00002741_a0/latents/restored/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..3378525fcb9ff125101c558fb4a0dfb1d55d5e2d --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/restored/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fbe503150e57f408a6238d21ca5f88d88f495d8be137306f1074fcb01976192d +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/restored/other.pt b/demo_samples/val_cache/00002741_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..dfc6d5c5f7f5ea7f1db1f0d1c965261ea6ac389b --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ff905458e8bff35d5cf22341b02846097ce9d9cacba8f8f4e783a53f98daefcb +size 18395 diff --git a/demo_samples/val_cache/00002741_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00002741_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..b55f5579ac57a688d89ee69ba63215eb5b048b9b --- /dev/null +++ b/demo_samples/val_cache/00002741_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a19a2231b1e194522da6ff3463ae6ebafd40d84972b9c7960f696ad8b09ba35 +size 18402 diff --git a/demo_samples/val_cache/00002795_a0/latents/clean/other.pt b/demo_samples/val_cache/00002795_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..ac4af34bd9246726be6f058d56a8f7ec0a4bc0a2 --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ad1d1eb05309122c4336d498446a52bec18b9a4857a5bc582aec2c76e945e8f8 +size 18395 diff --git a/demo_samples/val_cache/00002795_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00002795_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..ec957e9caea402ce0e92e20e6d7e770892bf2dda --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1435a4c49433afc7b824ccab15bc3781b5e289b37487b190ac16a986f852f638 +size 18402 diff --git a/demo_samples/val_cache/00002795_a0/latents/degraded/other.pt b/demo_samples/val_cache/00002795_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..b68ae227410769951280cb1fcd6bd2e612dffb2d --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3b9b59a7b74bd830670f7276703a5cca58c4d96d65895bf4ae070fd40ecd50f +size 18395 diff --git a/demo_samples/val_cache/00002795_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00002795_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..d6a30bebae12eb1e6d03423fdd2a2b1ad5a7cc49 --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:597511f8d6e8c26865be937980fb487bd17d6db5d9f2684e337fc26b64fffb13 +size 18402 diff --git a/demo_samples/val_cache/00002795_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00002795_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..6b9ac9f96c4168a96c96d2bb1c5c914faad253b1 --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:12d45c7244ff448fdc3da5288c6b41204dd8ff65ea83498689ad764e0a578daf +size 18508 diff --git a/demo_samples/val_cache/00002795_a0/latents/restored/other.pt b/demo_samples/val_cache/00002795_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..a0cf72fb3961ed088243aabe6d4424fff3eadf38 --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2da9c95b54fe4725ea825727c69c6af1b4bc9f90cdfd05ac8eed1ed06d0aa0b8 +size 18395 diff --git a/demo_samples/val_cache/00002795_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00002795_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..22b9cb0bcae519d2d9d7dd40bbc747b384959b71 --- /dev/null +++ b/demo_samples/val_cache/00002795_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2ca353e47607826f89e30f02ff55a1f004a51fbc63ab217a7ec722a8f664267d +size 18402 diff --git a/demo_samples/val_cache/00002878_a0/latents/clean/drums.pt b/demo_samples/val_cache/00002878_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..e7cc2aa0a128863dbf46b138d168f58ec09bfcea --- /dev/null +++ b/demo_samples/val_cache/00002878_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bb664dfe276a9a4bda39dd999db427df5e5cf9bab82520bcf6ae99bd1aa87c38 +size 18395 diff --git a/demo_samples/val_cache/00002878_a0/latents/clean/other.pt b/demo_samples/val_cache/00002878_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..3cd0c51c6ec2ef101d5cb57c9f41cbd4c8149ad5 --- /dev/null +++ b/demo_samples/val_cache/00002878_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6f950dc64a2a86b0d8ef06f66fd334d1bb1660569efe150b881a989283b24b08 +size 18395 diff --git a/demo_samples/val_cache/00002878_a0/latents/degraded/other.pt b/demo_samples/val_cache/00002878_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..f34bda3953480d157349585528db06b5f8b2e70d --- /dev/null +++ b/demo_samples/val_cache/00002878_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c86927e9ae0ea10f7c444402db4a70a4970d183182dfd73cee6ef5376764a1dd +size 18395 diff --git a/demo_samples/val_cache/00002878_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00002878_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..677b7233582eb6f31ac0fb7b4c3655479af3b9ff --- /dev/null +++ b/demo_samples/val_cache/00002878_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23df813d979e666f4365dadf28b4ab73157a43bf4e41c6c64cd4b11210b7e1d8 +size 18508 diff --git a/demo_samples/val_cache/00002878_a0/latents/restored/other.pt b/demo_samples/val_cache/00002878_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..ba3dcb6b2fff597873076c4a024cf02113208be7 --- /dev/null +++ b/demo_samples/val_cache/00002878_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:db8a9d2860063c96515a3818d375c0439132f11c427ba7e0c37291606cf00c8f +size 18395 diff --git a/demo_samples/val_cache/00003027_a0/latents/clean/other.pt b/demo_samples/val_cache/00003027_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..4fd71ef45528f7035b5110fbb3925521251e46fc --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2e7151e3ac046f9e803100a2f6d6454cfb7721abbf6872c635b0c0bf42f02cd5 +size 18395 diff --git a/demo_samples/val_cache/00003027_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00003027_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..551ec62cdb7f2840291545c0ca4974b92631baa9 --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b17a75708bedb4c2a1ec19dbdc7fe915ab83d427755c05fbcb8f3610a610380 +size 18402 diff --git a/demo_samples/val_cache/00003027_a0/latents/degraded/other.pt b/demo_samples/val_cache/00003027_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..e90f7022c3399f333871afb5c2d1b8b00e96e350 --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:41e29233bb92149c66106c3857469d1acfb96093621e24e7313ac564259bd13e +size 18395 diff --git a/demo_samples/val_cache/00003027_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00003027_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..47d74fb18c3bf2996688381b74aad90e5fa1cb8a --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ee91497561cb8f5e896701c9972de90247be264cff0b4f5c2ae107e4894f5c63 +size 18402 diff --git a/demo_samples/val_cache/00003027_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00003027_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..f36f3df19bbb95f2ce5d42b920537fd38ca2fc4c --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:84a26ebc4d14eb30e41d2989cee4f35709d9a02ef13ef86f965eea47a474ee2d +size 18508 diff --git a/demo_samples/val_cache/00003027_a0/latents/restored/other.pt b/demo_samples/val_cache/00003027_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..0c689aff73717185ec5f42f55d71e342fc5dbfd0 --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76f25d9e79b93cf8f20c25ad565ad1ba3b68a5b0551195ada3ef6fa84bcdc599 +size 18395 diff --git a/demo_samples/val_cache/00003027_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00003027_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..9b4bcd3e56bec16db2929eee06827a8f5eae12f2 --- /dev/null +++ b/demo_samples/val_cache/00003027_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c0c12c74e3d15107bd8c8033e0c5cfc37e494faefc873181428217020cdbcc1f +size 18402 diff --git a/demo_samples/val_cache/00003738_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00003738_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..b73668e8a6f7cadb4da8e86173013c3aae182c82 --- /dev/null +++ b/demo_samples/val_cache/00003738_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f2687b3ab37b6065c5b5d214dc0cf84f2d9ade1ea3c2c249f7bc1d9a5a27f496 +size 18402 diff --git a/demo_samples/val_cache/00003738_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00003738_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..ee9a118a48b2ae7675c193bd4c953f682f099c92 --- /dev/null +++ b/demo_samples/val_cache/00003738_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eda0ab3ccb88086a1268107acb6150378de903dc25fb64cf41bace0902f7f262 +size 18402 diff --git a/demo_samples/val_cache/00003738_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00003738_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..539366a3134946310cba9593ddb0e7f071ce30bd --- /dev/null +++ b/demo_samples/val_cache/00003738_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aea9b6e35e91123e6a0bf5c1c3c4b8009b80e2a16bc1378a478a6462b5cefa48 +size 18508 diff --git a/demo_samples/val_cache/00003738_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00003738_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..34a9367a94ab564b0bffb6639fdf5fd3d08afc25 --- /dev/null +++ b/demo_samples/val_cache/00003738_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e498129b3caa569f3eb440470c945ef38295b10ce27c2cdcced477da655ef0f +size 18402 diff --git a/demo_samples/val_cache/00003839_a0/latents/clean/drums.pt b/demo_samples/val_cache/00003839_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..ead1d50750999864cfe7dcb715af28c6a1fecd40 --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e695ef4d9dd63ee8d3c30f0c445bceb19d26ddde6ff3bd8fc1577ce3783bc87e +size 18395 diff --git a/demo_samples/val_cache/00003839_a0/latents/clean/other.pt b/demo_samples/val_cache/00003839_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..9a0b58451e7be4eb9ca75f7f1fd9e12c9fdd3d0e --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e076bafba1e9185d84ab7304ef485d5846dc01437707b2483a30c992243f246 +size 18395 diff --git a/demo_samples/val_cache/00003839_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00003839_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..6689f280bfccac9d2e0fb520e2b37dde56ffa954 --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dcd7cceeb2590d9c38435611cc8cf071fe583ca132b9f63c19ee820437cf2faa +size 18402 diff --git a/demo_samples/val_cache/00003839_a0/latents/degraded/other.pt b/demo_samples/val_cache/00003839_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..55728500516b0a3808b833f45c232de3f6cfc8a7 --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8ac2edd63ac2a879aaee8e2368d568fd1cefca2a76484bcdd68e72680228c61 +size 18395 diff --git a/demo_samples/val_cache/00003839_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00003839_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..14f94a12e349c0f3af50cb73688c3c056abf2b3f --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c150f9eacb5de1217ecbc946b0904472d320a075653da4e874feeba205b54b2d +size 18508 diff --git a/demo_samples/val_cache/00003839_a0/latents/restored/other.pt b/demo_samples/val_cache/00003839_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..a66ffb8856338f413dbb623becb3b97b6dd94bf3 --- /dev/null +++ b/demo_samples/val_cache/00003839_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1d90e90f2a921df3ccb8d9a9d967b77753982ddeb26826525d5fb5fb686309d0 +size 18395 diff --git a/demo_samples/val_cache/00003914_a0/latents/clean/bass.pt b/demo_samples/val_cache/00003914_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..c59a5d375aa71bb7117de7ff3d4c339c75bcd91f --- /dev/null +++ b/demo_samples/val_cache/00003914_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09fa077aeeadbb63010c029d8f16005751bc7bf0728b5d1885d2907e3b226692 +size 18388 diff --git a/demo_samples/val_cache/00003914_a0/latents/clean/other.pt b/demo_samples/val_cache/00003914_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..c7e546cc181aefa3d2adef0cb0f2e1a40e5a304f --- /dev/null +++ b/demo_samples/val_cache/00003914_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8d8955c17e810fdf087e6685e2016aa8bc7b60fe578f6c19db53d38767a8ecb +size 18395 diff --git a/demo_samples/val_cache/00003914_a0/latents/degraded/other.pt b/demo_samples/val_cache/00003914_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..37cd2042b17e2dd1e896829fc6a42f9137da6e55 --- /dev/null +++ b/demo_samples/val_cache/00003914_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e2321271e7c81acafe53379df5c403984679927dab7ede2619f0cc0e369d6a0 +size 18395 diff --git a/demo_samples/val_cache/00003914_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00003914_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..700ab03e7b18e1c44b52aeeb69a659fa3001362f --- /dev/null +++ b/demo_samples/val_cache/00003914_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1db60de1533c790073287f388b29e81b9f825a7980ee9ff91b1ee3dfcbaf5881 +size 18508 diff --git a/demo_samples/val_cache/00003914_a0/latents/restored/other.pt b/demo_samples/val_cache/00003914_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..b3e8fd977d209a78a6083e9d956355f4eee25e60 --- /dev/null +++ b/demo_samples/val_cache/00003914_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a430f46bc45a7d247670a9a4f58a2e5dc95362b1d0ca927aa8fa1f9926002179 +size 18395 diff --git a/demo_samples/val_cache/00004113_a0/latents/clean/other.pt b/demo_samples/val_cache/00004113_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..c28686288843c9d36c354a7ee7ce365f378016b2 --- /dev/null +++ b/demo_samples/val_cache/00004113_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:862674e0d7debfd908d64fbf7c31a40db87acd183062227c9642e13382e5a92a +size 18395 diff --git a/demo_samples/val_cache/00004113_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004113_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..5a6d096438c7eec1a8069338977ca861928e50c9 --- /dev/null +++ b/demo_samples/val_cache/00004113_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4625a9e2d3ef6454f74e3f0f5ea12fcf0e11e36fee77ce86a116b07f4305bb35 +size 18402 diff --git a/demo_samples/val_cache/00004113_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004113_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..b9e46043d1b97a1e6b2ca093763e001abb5f6143 --- /dev/null +++ b/demo_samples/val_cache/00004113_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0b05df0a6c4d2d1ae19246e64ba582b5820d3ba957d6fc38ab2ed1942791acda +size 18395 diff --git a/demo_samples/val_cache/00004113_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004113_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..a5aba86d1e1face2ad620459f46b29a97089cdab --- /dev/null +++ b/demo_samples/val_cache/00004113_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e764ce4d11aa6989fb8afe62dcfe9890bad6a64d4d34122a71aed9a000f566d8 +size 18508 diff --git a/demo_samples/val_cache/00004113_a0/latents/restored/other.pt b/demo_samples/val_cache/00004113_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..a2255decfd034dcbd9fe971bfbb37fa0174adfef --- /dev/null +++ b/demo_samples/val_cache/00004113_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a7f08185429426c9fce4db94a0e07070458cc18780d25db86fefb9938c2123d4 +size 18395 diff --git a/demo_samples/val_cache/00004291_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004291_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..b2e62b1384445b0e62a022f8432c772e5ca4aa5f --- /dev/null +++ b/demo_samples/val_cache/00004291_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:19ec1f73d9682385d65d6f845fe566e406128b05389e87bbc89b3141c2dbd9ae +size 18402 diff --git a/demo_samples/val_cache/00004291_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00004291_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..265b132b619575ae8a11689c68ce60addacc311a --- /dev/null +++ b/demo_samples/val_cache/00004291_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d9f7f674fb18175d440eb8fd0376f8ff6bb28fc270210451098649641824b774 +size 18402 diff --git a/demo_samples/val_cache/00004291_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004291_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..d57c4f46e31dc1e91d328659b38a745890e946c0 --- /dev/null +++ b/demo_samples/val_cache/00004291_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:916cab97d8a573488507175f5876fe7238b7a4c7d7829298ec942e8efe1cc4b9 +size 18508 diff --git a/demo_samples/val_cache/00004291_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00004291_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..ffb79399c79763866e523077ca4d823d35b6a4d5 --- /dev/null +++ b/demo_samples/val_cache/00004291_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b37572022a5295dcccc6f04c0e7468059b854d29c4a8a3215e79abe15ba9e53 +size 18402 diff --git a/demo_samples/val_cache/00004450_a0/latents/clean/other.pt b/demo_samples/val_cache/00004450_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..250b710e05415e52f461cbe4ee03ee237bf17369 --- /dev/null +++ b/demo_samples/val_cache/00004450_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a8c7e80c171eb998913b4d166180936b1eac8654977831b3847f4dd12f810e46 +size 18395 diff --git a/demo_samples/val_cache/00004450_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004450_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..751b4dce855a8b889372114a2118c93260b6f3fd --- /dev/null +++ b/demo_samples/val_cache/00004450_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2cc6fa420f46abbeed6f3d53cc006d0d983b70d63254627196ab1702c9f00518 +size 18402 diff --git a/demo_samples/val_cache/00004450_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004450_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..10cd29a332dce0fe567d8777d8cc9d6c5961c477 --- /dev/null +++ b/demo_samples/val_cache/00004450_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09cc1367831d928bf28dcdd05803dd75e211766015ed669c5977f69d98466133 +size 18395 diff --git a/demo_samples/val_cache/00004450_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004450_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..07ad6f14146a87cb511d45466a0933b6ba845bcf --- /dev/null +++ b/demo_samples/val_cache/00004450_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20e11a36ddea3a77943c92dc32102e4c474f42c584b1275d8e6914d46e0faa62 +size 18508 diff --git a/demo_samples/val_cache/00004450_a0/latents/restored/other.pt b/demo_samples/val_cache/00004450_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..a28f77ee7baf02a134a4031a56800cb83bb9511a --- /dev/null +++ b/demo_samples/val_cache/00004450_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:60d35d21ccfe5b2e459e947bf57bf62dfeb2c108da6dd0d0e54dce34a2538af5 +size 18395 diff --git a/demo_samples/val_cache/00004480_a0/latents/clean/drums.pt b/demo_samples/val_cache/00004480_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..e0255e7cace924270a012269a7056e490672b4ff --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b936c3ece358e6ef91969d5004f5fdb696da201535a08ca6c70f9726359fc181 +size 18395 diff --git a/demo_samples/val_cache/00004480_a0/latents/clean/other.pt b/demo_samples/val_cache/00004480_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..fa770b5229045150c64d9e44d07206a99ee99840 --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:00284c85592d50849cd5024cb44bcca9dadc01356b4132e7509eb746251198e0 +size 18395 diff --git a/demo_samples/val_cache/00004480_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004480_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..c9e48b0606cf42f5a4c8f28665dfee47064efe9f --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37035ca1fd0950d289d5fecf0330eade6a82d4ac364ce23be65af2b7e70e3ff0 +size 18402 diff --git a/demo_samples/val_cache/00004480_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004480_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..bc0c56ce71dd3ff65356a8da5fcc377de4adca93 --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ae8335f59f1d1412fcde9d616f46099aaedd2c2f5f9145e467b6f3acbf60065 +size 18395 diff --git a/demo_samples/val_cache/00004480_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00004480_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..536764544033f0edad22e9d7f5352945d335ac1a --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fea5cf0f1dfe471210e6ffc7ba0556c7e6011d0fe533256a40a0fd79bacfb45a +size 18402 diff --git a/demo_samples/val_cache/00004480_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004480_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..6b179897a8d1a37f44d44180b50ddf11c8e0719f --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:54737c056691182714f16d90ba871283fde9febabf1ab6947e01fd0cf7f66350 +size 18508 diff --git a/demo_samples/val_cache/00004480_a0/latents/restored/other.pt b/demo_samples/val_cache/00004480_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..1cd497bcc0838362baf6b098dc40be6efba09248 --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76d40b739f368e7560ddb023ba4133c20cb97133e7490bb503878f34c7c2c3b5 +size 18395 diff --git a/demo_samples/val_cache/00004480_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00004480_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..d78c80c3947ce19347bd222ec2c5d8f6830ffc64 --- /dev/null +++ b/demo_samples/val_cache/00004480_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:cc24d1408b8ea932465b40affa5ac81b3896a02979773f51b42ff55862fda5d7 +size 18402 diff --git a/demo_samples/val_cache/00004656_a0/latents/clean/bass.pt b/demo_samples/val_cache/00004656_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..7ba6b9feed3d69abafaf0f78afd1a8f05867fcbb --- /dev/null +++ b/demo_samples/val_cache/00004656_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:23608ae1bd5df09fff6369c03d2739de2bc09b87d68f14af7d929e311bc4ba7a +size 18388 diff --git a/demo_samples/val_cache/00004656_a0/latents/clean/other.pt b/demo_samples/val_cache/00004656_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..6af22d6f944c2f2aad5779234add448581bb76ba --- /dev/null +++ b/demo_samples/val_cache/00004656_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1893a70a39a7fb876a1d4ea5979c8e6c6e18ea73a4dc5950f712ea50307f071c +size 18395 diff --git a/demo_samples/val_cache/00004656_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004656_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..750aabe98dfdaf11fcda35852d155f6ad236d616 --- /dev/null +++ b/demo_samples/val_cache/00004656_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:81a16840f682f5a6aad3d429b4f24e256130e016c2ea1bcfaa5b122df2c4abdb +size 18395 diff --git a/demo_samples/val_cache/00004656_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004656_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..49e99318e5ef54953c9dec3d6d33c9fabda9a043 --- /dev/null +++ b/demo_samples/val_cache/00004656_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2d8dd15e2fc630b59427870804b1162b78b96be563790e8fa4c34204f37a307f +size 18508 diff --git a/demo_samples/val_cache/00004656_a0/latents/restored/other.pt b/demo_samples/val_cache/00004656_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..dbc2d110cd63d16b09329d1c4ee09129dd2f741b --- /dev/null +++ b/demo_samples/val_cache/00004656_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e2bc501543a6161fab3d2e2fbd729862805b073a0fd302efe57f331ae9914c5f +size 18395 diff --git a/demo_samples/val_cache/00004788_a0/latents/clean/other.pt b/demo_samples/val_cache/00004788_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..1a72ba804d65b5d64c1e5d6f86fc1d2f5150cd1d --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a68aa891f6e52afd477e8b1ec578cb8138db5108568fd50631eb47652687603c +size 18395 diff --git a/demo_samples/val_cache/00004788_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004788_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..3e6bbf904e04cd549b52591b39b63abba321dd6b --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ab6d6ad834c2d9539edb6edd813a3d43948a2eedfc4f5a6cdf5f867066b4fc2c +size 18402 diff --git a/demo_samples/val_cache/00004788_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004788_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..0f12a47950935db596c7960d03e1ba286fee0106 --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3772d30cfdf15922d2f2f33baf87f1b030cd0014f82b6b6313ad71aad83db33d +size 18395 diff --git a/demo_samples/val_cache/00004788_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00004788_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..d91506a0648d7b2b2f0bb1a993709d3ee9f2a423 --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:77a10b3d2f8420f79b3c517f9b853fb7e08538c655f9e1645dbb1a05b8b181b8 +size 18402 diff --git a/demo_samples/val_cache/00004788_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004788_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..e973cb774f8ef35ffe7176c96bd8a553eb3dacf0 --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:626576704a272c8842bf2b33031752912ff24fd125afb3345973da4734ceb3fc +size 18508 diff --git a/demo_samples/val_cache/00004788_a0/latents/restored/other.pt b/demo_samples/val_cache/00004788_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..15cb81a43808fe53b33cee4a12189a275d3420b8 --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e5e96d6e633ef0b7eee5999b4e4bde9ce7153854b6ded5fd447a37d0a6b2588c +size 18395 diff --git a/demo_samples/val_cache/00004788_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00004788_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..91e184b9bbe40fed0c4a5ec6b9d5f8e285c324fe --- /dev/null +++ b/demo_samples/val_cache/00004788_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:731f9996a522d09930831da41f3a7b0fca08133b9219fa255de2ef8b1584c79c +size 18402 diff --git a/demo_samples/val_cache/00004822_a0/latents/clean/other.pt b/demo_samples/val_cache/00004822_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..7972545b55281213868e7586dc7cefbab9340229 --- /dev/null +++ b/demo_samples/val_cache/00004822_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9239dc1e60d7dea0ff196923fae90147f2c43926cf18533993263f9fea008584 +size 18395 diff --git a/demo_samples/val_cache/00004822_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00004822_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..7a51bd8d2794d86e37105658968d218f47666bf4 --- /dev/null +++ b/demo_samples/val_cache/00004822_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3e70caeb00e1e54a07f8ad9de482bcf0c987bac10c693fa33596ea311e39d0d +size 18402 diff --git a/demo_samples/val_cache/00004822_a0/latents/degraded/other.pt b/demo_samples/val_cache/00004822_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..16c73d2ac29034c0d2c6a4afa064a7aaabdd5eb9 --- /dev/null +++ b/demo_samples/val_cache/00004822_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ccc359964a78ffe24f7eb2395a60ef83f5cc8ac2337d27ddfbf3da15926d81cc +size 18395 diff --git a/demo_samples/val_cache/00004822_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00004822_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..d34be353afb398af442a8fb359d67696b2faec37 --- /dev/null +++ b/demo_samples/val_cache/00004822_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:20bc2dd3009f530adba127e5ce3254956dad85d29d38d3afdfe1480ce99bf4be +size 18508 diff --git a/demo_samples/val_cache/00004822_a0/latents/restored/other.pt b/demo_samples/val_cache/00004822_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..c9e74711e45c5ffdfffedb63249da23bbd57f2da --- /dev/null +++ b/demo_samples/val_cache/00004822_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa69ab43b7c3657b95ed4c2f1d1bc65991fc02cf4264d97c6ddecedd2af2ef54 +size 18395 diff --git a/demo_samples/val_cache/00005806_a0/latents/clean/bass.pt b/demo_samples/val_cache/00005806_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..11c7bd12b2a2db7e6aa9cb72a205f72aaef15f55 --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b86e054e4b7b882af1a2fff471d7b0c45555eab1064d90ab91434ffa65332a1 +size 18388 diff --git a/demo_samples/val_cache/00005806_a0/latents/clean/drums.pt b/demo_samples/val_cache/00005806_a0/latents/clean/drums.pt new file mode 100644 index 0000000000000000000000000000000000000000..a296df0928f7f9300eedce7f9b1d12616f444b5b --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/clean/drums.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4655a576348ba6c35acecc50e831c9d27e9f6f1fc92847b159475a22dad0d177 +size 18395 diff --git a/demo_samples/val_cache/00005806_a0/latents/clean/other.pt b/demo_samples/val_cache/00005806_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..3cb391b5281eee4ef570c331f41f30955d30ed76 --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a38a9cd145c82ca0375f1951acdcc2fa5a3c411804ad2d46a9b8dc5382e1f631 +size 18395 diff --git a/demo_samples/val_cache/00005806_a0/latents/degraded/other.pt b/demo_samples/val_cache/00005806_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..fbfe1ae56ba04cb060c30311a4417f1ede3e094e --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3752617743facabb95cac2d4a0fafc51a02b24410a0cb6a59674381e1486a275 +size 18395 diff --git a/demo_samples/val_cache/00005806_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00005806_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..6a4e559074b3f79fa467742dcdf97cdede00966c --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:99285067c571542422552fac4a070555ed12d8d0ea3631d06b75581c6d2d07e4 +size 18508 diff --git a/demo_samples/val_cache/00005806_a0/latents/restored/other.pt b/demo_samples/val_cache/00005806_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..1891c44fd7d7436d49757db456e542315aee93e9 --- /dev/null +++ b/demo_samples/val_cache/00005806_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0f07ab4bee2d828c1bcb94dda2bf0618b794dc2146f2cb68207b5fccd5bb38bf +size 18395 diff --git a/demo_samples/val_cache/00005939_a0/latents/clean/other.pt b/demo_samples/val_cache/00005939_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..5f1476bd65c1f08800d9c7e064768f44323dcf77 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:45485fe129dc00ebd070a34eb50b9c7623c163339e9f9733b57e673f2f52b066 +size 18395 diff --git a/demo_samples/val_cache/00005939_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00005939_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..ede3e982b14c992aad4143108e87ef9e07ba7822 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b9ecb09bc230640b202da96492ef311c18f54843bf0e055199d0effc7590572f +size 18402 diff --git a/demo_samples/val_cache/00005939_a0/latents/degraded/other.pt b/demo_samples/val_cache/00005939_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..c51c652711980d04cf6f3c304cefe439c17d681a --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3b5b38d20cf90bac7514318bf958f3e103bafb1512b2be9a92f056aa1fea8aad +size 18395 diff --git a/demo_samples/val_cache/00005939_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00005939_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..b5ca44da791dc35a0be744f2de849c5a8f502b10 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1c661b3135241053e33089c7e6ad9a5044bd71d11dd3677616ddcbff0235e0e1 +size 18402 diff --git a/demo_samples/val_cache/00005939_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00005939_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..7138953e9e1f9f85d931739afa39926c72620808 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d2f37e87782bd2a315ea450a6d454d715fd60ac1cc6f362d196fcdf7ef0408dd +size 18508 diff --git a/demo_samples/val_cache/00005939_a0/latents/restored/other.pt b/demo_samples/val_cache/00005939_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..d0cd59438b21a0abc413d62b9cff1429b986b2b8 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76ea55b23c22c54fc18df3734765bcca90a9902a79182704d973cfda848ce446 +size 18395 diff --git a/demo_samples/val_cache/00005939_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00005939_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..5020d4c9b1189a87268e2f486ffe0b7a37fd83b8 --- /dev/null +++ b/demo_samples/val_cache/00005939_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5797e1b6576cfa97a6c49d0826e31f37d63126e7661807e7b6cd2f6ce3240ddb +size 18402 diff --git a/demo_samples/val_cache/00006037_a0/latents/clean/bass.pt b/demo_samples/val_cache/00006037_a0/latents/clean/bass.pt new file mode 100644 index 0000000000000000000000000000000000000000..b20646fc983e0187362612e2d57799cc15036d27 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/clean/bass.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:235b8d1f93d9bade36a8e5787917694192a57449ce4b20e72905706390aaf24f +size 18388 diff --git a/demo_samples/val_cache/00006037_a0/latents/clean/other.pt b/demo_samples/val_cache/00006037_a0/latents/clean/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..d95db17b794d6d80d2de0be056093bca3ee9c428 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/clean/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3c3377bb1c8bcad5ad6412da5f21b2474d282798bae4478677a0b8afd18d20ee +size 18395 diff --git a/demo_samples/val_cache/00006037_a0/latents/clean/vocals.pt b/demo_samples/val_cache/00006037_a0/latents/clean/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..b4b67cfac0e0e0eb4f00e04911f0134da8fc2be3 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/clean/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5e74e3782e6b53f9a5087c863ee7340ed7b97581c7c676a08296ff3dd890152f +size 18402 diff --git a/demo_samples/val_cache/00006037_a0/latents/degraded/other.pt b/demo_samples/val_cache/00006037_a0/latents/degraded/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..62b017ac79ea08cc62a3b7974ba20cd81f5509a4 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/degraded/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f59d0bb02c916b3925215271270ef8d4470ed1498f34d9621092ff18026b9646 +size 18395 diff --git a/demo_samples/val_cache/00006037_a0/latents/degraded/vocals.pt b/demo_samples/val_cache/00006037_a0/latents/degraded/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..5a5654bd3d0d68a9c86a95a2298c282f04de7d6a --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/degraded/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:67cc1a1a8262638cdb0f2821931941ba8d0fb7856921b8581fce7efa889a14e9 +size 18402 diff --git a/demo_samples/val_cache/00006037_a0/latents/degraded_mix.pt b/demo_samples/val_cache/00006037_a0/latents/degraded_mix.pt new file mode 100644 index 0000000000000000000000000000000000000000..4c8f25d347ff3c77d35d090afbb68eee4a9152cf --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/degraded_mix.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dd7d1e964e6ef6049cd60acc0f34d668d1819f3f89bbca84ca6b0b27bb9fc494 +size 18508 diff --git a/demo_samples/val_cache/00006037_a0/latents/restored/other.pt b/demo_samples/val_cache/00006037_a0/latents/restored/other.pt new file mode 100644 index 0000000000000000000000000000000000000000..f470591e87017bc21f922ad0eeef16606088aee3 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/restored/other.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f68d3709a8061642058c0bfa5761ce30b492762712aa27931b229c632cbf5584 +size 18395 diff --git a/demo_samples/val_cache/00006037_a0/latents/restored/vocals.pt b/demo_samples/val_cache/00006037_a0/latents/restored/vocals.pt new file mode 100644 index 0000000000000000000000000000000000000000..551a41cd24257ecc5c0a2881a05702be32cef273 --- /dev/null +++ b/demo_samples/val_cache/00006037_a0/latents/restored/vocals.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6a6f72ff2e87ab9610f3c453cf11791512d3746b6157ec068bd15ee24fb072cb +size 18402 diff --git a/demo_samples/val_cache/metadata.csv b/demo_samples/val_cache/metadata.csv new file mode 100644 index 0000000000000000000000000000000000000000..d39720da0d926ac7bdb2cc44908840b99ac844ea --- /dev/null +++ b/demo_samples/val_cache/metadata.csv @@ -0,0 +1,97 @@ +pair_id,stem,class,clean_present,deg_present,clean_rms_db,deg_rms_db,clean_peak_db,deg_peak_db,clean_latent_path,deg_latent_path,latent_dim,n_frames,sample_rate,frame_rate,duration_s +00000486_a0,bass,both_silent,0,0,-73.21,-75.01,-56.87,-57.05,,,,,44100,,3.0 +00000486_a0,drums,both_silent,0,0,-73.37,-72.8,-54.75,-55.19,,,,,44100,,3.0 +00000486_a0,other,restoration,1,1,-15.31,-15.02,-0.91,-0.91,00000486_a0/latents/clean/other.pt,00000486_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00000486_a0,vocals,deg_only,0,1,-67.1,-26.76,-48.23,-7.92,,,,,44100,,3.0 +00000607_a0,bass,generation,1,0,-22.61,-69.71,-9.51,-53.66,00000607_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00000607_a0,drums,generation,1,0,-44.67,-68.53,-16.76,-43.86,00000607_a0/latents/clean/drums.pt,,256,33,44100,11.0,3.0 +00000607_a0,other,restoration,1,1,-17.53,-19.26,-4.04,-6.32,00000607_a0/latents/clean/other.pt,00000607_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00000607_a0,vocals,restoration,1,1,-13.3,-12.38,-3.69,-2.21,00000607_a0/latents/clean/vocals.pt,00000607_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00001171_a0,bass,restoration,1,1,-22.37,-27.38,-6.41,-11.83,00001171_a0/latents/clean/bass.pt,00001171_a0/latents/degraded/bass.pt,256,33,44100,11.0,3.0 +00001171_a0,drums,restoration,1,1,-33.52,-31.92,-7.12,-7.74,00001171_a0/latents/clean/drums.pt,00001171_a0/latents/degraded/drums.pt,256,33,44100,11.0,3.0 +00001171_a0,other,restoration,1,1,-12.57,-13.14,-0.91,-0.91,00001171_a0/latents/clean/other.pt,00001171_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00001171_a0,vocals,both_silent,0,0,-70.01,-71.3,-49.32,-56.33,,,,,44100,,3.0 +00001288_a0,bass,both_silent,0,0,-61.74,-57.1,-44.2,-39.83,,,,,44100,,3.0 +00001288_a0,drums,restoration,1,1,-23.61,-37.11,-5.9,-19.81,00001288_a0/latents/clean/drums.pt,00001288_a0/latents/degraded/drums.pt,256,33,44100,11.0,3.0 +00001288_a0,other,restoration,1,1,-14.17,-14.39,-0.92,-0.91,00001288_a0/latents/clean/other.pt,00001288_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00001288_a0,vocals,both_silent,0,0,-70.26,-68.69,-51.42,-53.28,,,,,44100,,3.0 +00001996_a0,bass,both_silent,0,0,-64.25,-77.47,-50.14,-59.68,,,,,44100,,3.0 +00001996_a0,drums,both_silent,0,0,-68.32,-73.14,-53.16,-51.72,,,,,44100,,3.0 +00001996_a0,other,restoration,1,1,-22.21,-27.21,-3.21,-6.42,00001996_a0/latents/clean/other.pt,00001996_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00001996_a0,vocals,restoration,1,1,-15.95,-17.7,-1.5,-0.92,00001996_a0/latents/clean/vocals.pt,00001996_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00002166_a0,bass,generation,1,1,-22.12,-40.61,-6.89,-20.86,00002166_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00002166_a0,drums,both_silent,0,0,-63.86,-66.63,-47.7,-50.31,,,,,44100,,3.0 +00002166_a0,other,restoration,1,1,-21.46,-19.1,-8.62,-6.06,00002166_a0/latents/clean/other.pt,00002166_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00002166_a0,vocals,restoration,1,1,-14.52,-15.12,-2.1,-1.54,00002166_a0/latents/clean/vocals.pt,00002166_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00002471_a0,bass,both_silent,0,0,-63.31,-74.25,-44.46,-50.05,,,,,44100,,3.0 +00002471_a0,drums,restoration,1,1,-28.69,-41.14,-6.89,-14.66,00002471_a0/latents/clean/drums.pt,00002471_a0/latents/degraded/drums.pt,256,33,44100,11.0,3.0 +00002471_a0,other,restoration,1,1,-13.11,-15.3,-0.92,-0.91,00002471_a0/latents/clean/other.pt,00002471_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00002471_a0,vocals,both_silent,0,0,-67.09,-71.38,-37.7,-55.35,,,,,44100,,3.0 +00002741_a0,bass,both_silent,0,0,-65.16,-68.05,-48.37,-52.03,,,,,44100,,3.0 +00002741_a0,drums,restoration,1,1,-23.22,-29.02,-5.06,-11.27,00002741_a0/latents/clean/drums.pt,00002741_a0/latents/degraded/drums.pt,256,33,44100,11.0,3.0 +00002741_a0,other,restoration,1,1,-21.7,-22.52,-4.6,-7.37,00002741_a0/latents/clean/other.pt,00002741_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00002741_a0,vocals,restoration,1,1,-15.18,-15.07,-1.74,-1.43,00002741_a0/latents/clean/vocals.pt,00002741_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00002795_a0,bass,both_silent,0,0,-71.67,-75.53,-54.75,-58.27,,,,,44100,,3.0 +00002795_a0,drums,both_silent,0,0,-71.62,-72.89,-53.28,-55.35,,,,,44100,,3.0 +00002795_a0,other,restoration,1,1,-26.35,-28.37,-6.04,-6.46,00002795_a0/latents/clean/other.pt,00002795_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00002795_a0,vocals,restoration,1,1,-17.2,-17.28,-2.81,-1.02,00002795_a0/latents/clean/vocals.pt,00002795_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00002878_a0,bass,both_silent,0,0,-62.04,-71.99,-47.39,-57.24,,,,,44100,,3.0 +00002878_a0,drums,generation,1,0,-38.75,-69.67,-14.31,-51.32,00002878_a0/latents/clean/drums.pt,,256,33,44100,11.0,3.0 +00002878_a0,other,restoration,1,1,-12.1,-13.55,-0.91,-0.91,00002878_a0/latents/clean/other.pt,00002878_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00002878_a0,vocals,both_silent,0,0,-56.57,-71.65,-35.71,-57.24,,,,,44100,,3.0 +00003027_a0,bass,both_silent,0,0,-61.1,-75.38,-46.96,-59.18,,,,,44100,,3.0 +00003027_a0,drums,both_silent,0,0,-67.3,-72.83,-51.72,-53.79,,,,,44100,,3.0 +00003027_a0,other,restoration,1,1,-18.51,-26.12,-5.67,-11.51,00003027_a0/latents/clean/other.pt,00003027_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00003027_a0,vocals,restoration,1,1,-13.8,-15.84,-0.91,-0.92,00003027_a0/latents/clean/vocals.pt,00003027_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00003738_a0,bass,both_silent,0,0,-69.96,-75.42,-51.32,-58.27,,,,,44100,,3.0 +00003738_a0,drums,both_silent,0,0,-59.66,-72.81,-38.14,-55.5,,,,,44100,,3.0 +00003738_a0,other,both_silent,0,0,-59.43,-47.57,-39.63,-25.43,,,,,44100,,3.0 +00003738_a0,vocals,restoration,1,1,-14.35,-12.89,-0.92,-0.92,00003738_a0/latents/clean/vocals.pt,00003738_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00003839_a0,bass,both_silent,0,0,-65.33,-74.48,-52.14,-57.84,,,,,44100,,3.0 +00003839_a0,drums,generation,1,0,-45.53,-67.92,-16.25,-47.14,00003839_a0/latents/clean/drums.pt,,256,33,44100,11.0,3.0 +00003839_a0,other,restoration,1,1,-12.0,-14.05,-0.91,-0.92,00003839_a0/latents/clean/other.pt,00003839_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00003839_a0,vocals,generation,1,0,-24.58,-71.05,-4.23,-56.16,00003839_a0/latents/clean/vocals.pt,,256,33,44100,11.0,3.0 +00003914_a0,bass,generation,1,0,-19.23,-50.74,-9.21,-34.24,00003914_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00003914_a0,drums,both_silent,0,0,-66.14,-73.22,-48.37,-56.33,,,,,44100,,3.0 +00003914_a0,other,restoration,1,1,-15.71,-15.81,-0.91,-0.91,00003914_a0/latents/clean/other.pt,00003914_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00003914_a0,vocals,both_silent,0,0,-70.09,-69.31,-50.14,-49.25,,,,,44100,,3.0 +00004113_a0,bass,both_silent,0,0,-61.58,-72.97,-44.92,-57.84,,,,,44100,,3.0 +00004113_a0,drums,both_silent,0,0,-66.59,-70.08,-49.8,-53.04,,,,,44100,,3.0 +00004113_a0,other,restoration,1,1,-12.88,-11.93,-0.92,-0.91,00004113_a0/latents/clean/other.pt,00004113_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004113_a0,vocals,generation,1,0,-26.36,-61.55,-7.48,-44.6,00004113_a0/latents/clean/vocals.pt,,256,33,44100,11.0,3.0 +00004291_a0,bass,both_silent,0,0,-75.12,-73.96,-58.71,-58.94,,,,,44100,,3.0 +00004291_a0,drums,both_silent,0,0,-73.85,-73.66,-57.44,-57.84,,,,,44100,,3.0 +00004291_a0,other,both_silent,0,0,-60.9,-60.79,-46.39,-46.28,,,,,44100,,3.0 +00004291_a0,vocals,restoration,1,1,-12.04,-12.27,-0.91,-0.91,00004291_a0/latents/clean/vocals.pt,00004291_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00004450_a0,bass,both_silent,0,0,-62.54,-75.34,-48.8,-59.43,,,,,44100,,3.0 +00004450_a0,drums,both_silent,0,0,-66.43,-71.11,-51.13,-54.46,,,,,44100,,3.0 +00004450_a0,other,restoration,1,1,-12.14,-14.21,-0.91,-0.91,00004450_a0/latents/clean/other.pt,00004450_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004450_a0,vocals,generation,1,0,-34.64,-64.81,-12.79,-50.22,00004450_a0/latents/clean/vocals.pt,,256,33,44100,11.0,3.0 +00004480_a0,bass,both_silent,0,0,-62.8,-74.9,-41.62,-56.68,,,,,44100,,3.0 +00004480_a0,drums,generation,1,0,-35.02,-69.94,-11.24,-48.3,00004480_a0/latents/clean/drums.pt,,256,33,44100,11.0,3.0 +00004480_a0,other,restoration,1,1,-21.46,-24.94,-3.45,-7.23,00004480_a0/latents/clean/other.pt,00004480_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004480_a0,vocals,restoration,1,1,-15.11,-16.49,-0.92,-2.12,00004480_a0/latents/clean/vocals.pt,00004480_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00004656_a0,bass,generation,1,0,-22.67,-72.6,-12.92,-54.89,00004656_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00004656_a0,drums,both_silent,0,0,-62.52,-70.42,-47.7,-54.19,,,,,44100,,3.0 +00004656_a0,other,restoration,1,1,-12.84,-15.23,-0.92,-0.92,00004656_a0/latents/clean/other.pt,00004656_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004656_a0,vocals,both_silent,0,0,-63.75,-70.17,-35.09,-46.12,,,,,44100,,3.0 +00004788_a0,bass,both_silent,0,0,-58.63,-74.03,-41.0,-57.44,,,,,44100,,3.0 +00004788_a0,drums,both_silent,0,0,-65.01,-70.79,-45.96,-51.72,,,,,44100,,3.0 +00004788_a0,other,restoration,1,1,-17.18,-21.35,-1.68,-3.89,00004788_a0/latents/clean/other.pt,00004788_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004788_a0,vocals,restoration,1,1,-14.73,-15.16,-0.92,-0.92,00004788_a0/latents/clean/vocals.pt,00004788_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00004822_a0,bass,both_silent,0,0,-62.8,-72.89,-49.1,-57.64,,,,,44100,,3.0 +00004822_a0,drums,both_silent,0,0,-66.15,-67.98,-50.85,-54.05,,,,,44100,,3.0 +00004822_a0,other,restoration,1,1,-13.84,-13.68,-0.91,-0.91,00004822_a0/latents/clean/other.pt,00004822_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00004822_a0,vocals,generation,1,0,-26.28,-68.34,-10.79,-52.14,00004822_a0/latents/clean/vocals.pt,,256,33,44100,11.0,3.0 +00005806_a0,bass,generation,1,0,-26.63,-73.89,-11.27,-57.05,00005806_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00005806_a0,drums,generation,1,0,-22.49,-69.46,-4.6,-50.22,00005806_a0/latents/clean/drums.pt,,256,33,44100,11.0,3.0 +00005806_a0,other,restoration,1,1,-14.19,-13.44,-1.62,-0.91,00005806_a0/latents/clean/other.pt,00005806_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00005806_a0,vocals,both_silent,0,0,-50.8,-66.04,-33.49,-47.57,,,,,44100,,3.0 +00005939_a0,bass,both_silent,0,0,-66.74,-70.98,-51.03,-55.99,,,,,44100,,3.0 +00005939_a0,drums,both_silent,0,0,-65.28,-68.84,-39.26,-51.82,,,,,44100,,3.0 +00005939_a0,other,restoration,1,1,-19.9,-20.45,-3.37,-5.31,00005939_a0/latents/clean/other.pt,00005939_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00005939_a0,vocals,restoration,1,1,-15.21,-15.95,-2.44,-2.9,00005939_a0/latents/clean/vocals.pt,00005939_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 +00006037_a0,bass,generation,1,0,-25.48,-66.28,-9.68,-44.2,00006037_a0/latents/clean/bass.pt,,256,33,44100,11.0,3.0 +00006037_a0,drums,both_silent,0,0,-69.08,-69.72,-52.14,-52.03,,,,,44100,,3.0 +00006037_a0,other,restoration,1,1,-26.27,-25.32,-13.36,-11.69,00006037_a0/latents/clean/other.pt,00006037_a0/latents/degraded/other.pt,256,33,44100,11.0,3.0 +00006037_a0,vocals,restoration,1,1,-13.25,-13.82,-0.91,-0.91,00006037_a0/latents/clean/vocals.pt,00006037_a0/latents/degraded/vocals.pt,256,33,44100,11.0,3.0 diff --git a/packages.txt b/packages.txt new file mode 100644 index 0000000000000000000000000000000000000000..20645e641240cb419f5fc66c14c1447e91daf669 --- /dev/null +++ b/packages.txt @@ -0,0 +1 @@ +ffmpeg diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000000000000000000000000000000000000..f470ab69bb605be84688f58b3ea769acc97936cb --- /dev/null +++ b/requirements.txt @@ -0,0 +1,10 @@ +--extra-index-url https://download.pytorch.org/whl/cpu +torch==2.7.1 +torchaudio==2.7.1 +numpy==1.26.4 +stable-audio-tools==0.0.20 +demucs==4.0.1 +soundfile +librosa +matplotlib +pytorch-lightning diff --git a/restoflow/__init__.py b/restoflow/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..0c0b4f0982a3a94b22cd0cfaef6046b461efa8a7 --- /dev/null +++ b/restoflow/__init__.py @@ -0,0 +1 @@ +"""SAME-latent restoration: presence router + conditional flow-matching restorer.""" diff --git a/restoflow/app.py b/restoflow/app.py new file mode 100644 index 0000000000000000000000000000000000000000..b01c9f6718d70e76b650bb49c18625111e2fcda1 --- /dev/null +++ b/restoflow/app.py @@ -0,0 +1,746 @@ +"""Gradio demo for the SAME-latent stem RESTORER (v4b) + GENERATOR pipeline. + +Unified pipeline: a 3s clip's degraded full-mix -> htdemucs stems -> SAME latents -> + - present stems : v4b deterministic RESTORE + - missing stems : conditional-flow GENERATE from the (restored) context +-> SAME decode -> remixed output. Val samples use the cached latents (fast); uploads run +the full htdemucs+SAME path and are chunked into 3s (last chunk padded with silence then cropped). + +Runs on CPU by default (set RESTOFLOW_DEVICE=cuda to use a GPU). +Launch: python -m restoflow.app +""" +from __future__ import annotations +import os, csv, glob, io +# Subprocess ffmpeg fix (audio-separator + librosa/audioread shell out to `ffmpeg`): +# a stale /home/maximos/miniconda3/lib on LD_LIBRARY_PATH shadows system libstdc++ and +# breaks the system ffmpeg; drop it and put a known-good conda ffmpeg first on PATH. +os.environ["LD_LIBRARY_PATH"] = ":".join( + p for p in os.environ.get("LD_LIBRARY_PATH", "").split(":") if p and "maximos/miniconda3" not in p) +_FFMPEG_DIR = "/home/ksoil/.conda/envs/ksoil_torch/bin" +if os.path.isdir(_FFMPEG_DIR): + os.environ["PATH"] = _FFMPEG_DIR + ":" + os.environ.get("PATH", "") +from pathlib import Path +import numpy as np +import torch +import soundfile as sf +import librosa +import librosa.display +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import gradio as gr + +STEM_EMOJI = {"other": "๐ŸŽน", "vocals": "๐ŸŽค", "drums": "๐Ÿฅ", "bass": "๐ŸŽธ"} + +from .config import Cfg, STEMS, STEM_ID +from .model import DetRestorer, CondFlow, AttnRestorer, AttnCondFlow +from . import eval as E +from . import gen as G + +# Paths are env-overridable so the same file runs locally and on a packaged HF Space. +# Local default = the dev tree; on HF the deploy bundle sets these (or we fall back to a +# repo-relative layout: /restoflow_runs, /demo_samples). +_LOCAL_BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") +_REPO_ROOT = Path(__file__).resolve().parent.parent # containing restoflow/ +BASE = Path(os.environ.get("RESTOFLOW_BASE", _LOCAL_BASE if _LOCAL_BASE.exists() else _REPO_ROOT)) +SRC = Path(os.environ.get( + "RESTOFLOW_SRC", + "/media/maindisk/melkor169/MSRKit/xlance-msr/stereo_a2sb/synthetic_out")) +if not SRC.exists(): # packaged Space: demo wavs live here + SRC = BASE / "demo_samples" +VAL_CACHE = Path(os.environ.get("RESTOFLOW_VAL_CACHE", str(BASE / "demucs_results_val_full"))) +if not VAL_CACHE.exists(): + VAL_CACHE = BASE / "demo_samples" / "val_cache" +DEVICE = os.environ.get("RESTOFLOW_DEVICE", "cpu") +T, D = 32, 256 +SR = 44100 +N_VAL = 80 +PRESENT_RMS, PRESENT_PEAK = -40.0, -25.0 # presence gate (rms OR peak), matches training +# Generation policy: every song has a low end, so a MISSING bass is always invented. +# A missing drums/other is legitimate (song may have none) -> we DON'T fabricate it; restore-only. +GEN_TARGETS = ("bass",) + +_M = {} # model cache + + +# ---------------- model loading ---------------- +def _disable_flash_attn(): + """SAME's transformer uses flash_attn (CUDA-only). On CPU, force the SDPA fallback + by nulling the module globals apply_attn() checks at call time.""" + import stable_audio_tools.models.transformer as _tf + for name in ("flash_attn_func", "flash_attn_kvpacked_func", "flash_attn_varlen_func", "index_first_axis"): + if hasattr(_tf, name): + setattr(_tf, name, None) + + +def _ckpt_path(run): + d = BASE / "restoflow_runs" / run + for f in ("ckpt_best.pt", "ckpt.pt"): + if (d / f).exists(): + return d / f + return None + + +def _default_runs(rests, gens): + """Preselected restorer/generator. Env override (set by the HF deploy) wins, else a + 'prefer newest best' heuristic, else first available.""" + dr = os.environ.get("RESTOFLOW_DEFAULT_REST") + dg = os.environ.get("RESTOFLOW_DEFAULT_GEN") + if dr not in rests: + # restorer_attn_advramp is the quality winner: w003 (dist-aux 6.89->5.56) -> w003_gan (GAN realism + # ->4.95) -> cotrain (joint co-train ->4.56) -> advramp (co-train w_adv 0.08: best drums 3.236 + + # vocals 3.519 + brightest spectra; REMIX 4.55 ~tied long30's 4.54). restorer_attn_w003_gan kept + # as a selectable A/B baseline. Fallbacks: w003_gan, v1. + dr = next((r for r in ("restorer_attn_advramp", "restorer_attn_w003_gan", + "restorer_attn_v1", "v4b_balonly") if r in rests), + rests[0] if rests else None) + if dg not in gens: + # gen_advramp_v1: retrained on advramp's restored conditioning -> FAD-REMIX 16.75->14.59 (-13%) + # vs gen_distvar_baseline (kept as the selectable A/B baseline). Both conv 10M (fast on CPU). + dg = next((g for g in ("gen_advramp_v1", "gen_distvar_baseline", "gen_v2", "gen_v3") if g in gens), + gens[0] if gens else "(none)") + return dr, dg + + +def list_runs(kind): + """Available run names (excluding obsolete) for kind in {'restorer','generator'}. + Convention: gen_* are generators; everything else is a restorer.""" + out = [] + for d in sorted(glob.glob(str(BASE / "restoflow_runs" / "*"))): + name = Path(d).name + # skip obsolete + scratch/experiment dirs (scout_*, _smoke_*, _scout*, remix_* GAN-refine + # runs) โ€” not deployable until validated and explicitly promoted + if (name == "obsolete" or name.startswith(("scout", "_", "remix", "exp")) + or not Path(d).is_dir() or _ckpt_path(name) is None): + continue + is_gen = name.startswith("gen") + if (kind == "generator") == is_gen: + out.append(name) + return out + + +def load_restorer(run): + ck = torch.load(_ckpt_path(run), map_location="cpu"); rc = ck["cfg"] + if rc.get("model_kind") == "attn": # Conformer-lite restorer + m = AttnRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS), + stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"], heads=rc.get("heads", 4)) + else: + m = DetRestorer(D, rc["hidden"], rc["depth"], n_stems=len(STEMS), + stem_emb=rc["stem_emb_dim"], use_mix=rc["use_mix"]) + m.load_state_dict(ck["model"]); m.eval().to(DEVICE) + stats = torch.load(BASE / "restoflow_runs" / run / "norm_stats.pt", map_location="cpu") + return dict(rest=m, rstats=stats, rest_run=run, rest_ep=ck.get("epoch"), + rest_params=sum(p.numel() for p in m.parameters()) / 1e6) + + +def load_generator(run): + if run in (None, "(none)"): + return dict(gen=None, gstats=None, gargs=None, gen_run="(none)", gen_ep=None, gen_params=0.0) + ck = torch.load(_ckpt_path(run), map_location="cpu"); ga = ck["args"] + if ga.get("arch") == "attn": # Conformer-lite velocity net + m = AttnCondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS), + stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8)) + else: + m = CondFlow(D, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim) + m.load_state_dict(ck["model"]); m.eval().to(DEVICE) + return dict(gen=m, gstats=ck["stats"], gargs=ga, gen_run=run, gen_ep=ck.get("epoch"), + gen_params=sum(p.numel() for p in m.parameters()) / 1e6) + + +def status_md(): + if "rest_run" not in _M: + return "โณ **Models load on first run** (SAME + restorer + generator) โ€” first action takes ~15 s, then it's fast." + rb = f"**Restorer:** `{_M.get('rest_run','-')}` (ep {_M.get('rest_ep','?')}, {_M.get('rest_params',0):.1f}M)" + g = _M.get("gen_run", "(none)") + gt = (_M.get("gargs") or {}).get("target_stems", "โ€”") + gb = f"**Generator:** `{g}`" + (f" (ep {_M.get('gen_ep','?')}, {_M.get('gen_params',0):.1f}M, targets={gt})" if _M.get("gen") is not None else " โ€” *disabled*") + return f"๐ŸŽ›๏ธ Active models ยท {rb} ยท {gb} ยท **Autoencoder:** `SAME-L` ยท device `{DEVICE}`" + + +def set_models(rest_run, gen_run): + _M.update(load_restorer(rest_run)); _M.update(load_generator(gen_run)) + print(f"[app] active: restorer={rest_run} generator={gen_run}") + return status_md() + + +def models(): + if _M: + return _M + cfg = Cfg(); cfg.device = DEVICE + if DEVICE == "cpu": + _disable_flash_attn() + sa, sr = E.load_same(cfg) # SAME-L autoencoder (frozen) + _M.update(sa=sa, sr=sr, cfg=cfg) + rests = list_runs("restorer"); gens = list_runs("generator") + default_rest, default_gen = _default_runs(rests, gens) + _M.update(load_restorer(default_rest)); _M.update(load_generator(default_gen)) + print(f"[app] models loaded on {DEVICE}. restorer={default_rest} generator={default_gen}") + return _M + + +# ---------------- core latent ops ---------------- +def _fitT(x): return G._fitT(x.float(), T) + +def rms_peak_db(audio_mono): + rms = float(np.sqrt(np.mean(audio_mono ** 2)) + 1e-12) + peak = float(np.max(np.abs(audio_mono)) + 1e-12) + return 20 * np.log10(rms), 20 * np.log10(peak) + +@torch.inference_mode() +def restore(latent, stem, mix): + m = models(); st = m["rstats"]; mu, sd = st[stem]["mu"], st[stem]["sd"] + mm = st["__mix__"] + sd_d = ((latent - mu[:, None]) / sd[:, None]).to(DEVICE)[None] + sd_m = ((mix - mm["mu"][:, None]) / mm["sd"][:, None]).to(DEVICE)[None] + sid = torch.tensor([STEM_ID[stem]], device=DEVICE) + out = m["rest"](sd_d, sid, sd_m)[0].cpu() + return out * sd[:, None] + mu[:, None] + +@torch.inference_mode() +def generate(context_latents, stem): + m = models() + if m["gen"] is None: + return None + gs = m["gstats"]; cm, cs = gs["ctx"][stem]; tm, ts = gs["tgt"][stem] + ctx_raw = sum(_fitT(c) for c in context_latents) + ctx = ((ctx_raw - cm[:, None]) / cs[:, None]).to(DEVICE)[None] + sid = torch.tensor([STEM_ID[stem]], device=DEVICE) + w = m["gargs"].get("cfg_w", 2.0); steps = m["gargs"].get("sample_steps", 40) + z = G.sample_gen(m["gen"], ctx, sid, steps, w)[0].cpu() + return z * ts[:, None] + tm[:, None] + +# STOPGAP loudness control (proper fix = the planned generator overhaul). Generated bass decodes +# ~2x too loud (off-manifold flow endpoints); neither more data nor attention fixed it. We ground +# its level in the context: target rms = ratio x restored-context rms, applied in AUDIO space +# (SAME decode is nonlinear). Ratio = geometric mean of clean_bass/ctx over GENERATION-case val +# pairs (0.24); the per-song spread is wide (p25 .12, p75 .81), so we deliberately err quiet โ€” +# too-loud bass is far more jarring than too-quiet. +GEN_ENERGY_RATIO = {"bass": 0.24} + +def energy_match(audio_by_stem: dict, dec: dict): + """Scale each GENERATED stem to its fitted energy relative to the restored context.""" + ctx = [audio_by_stem[s] for s, v in dec.items() if v == "restored" and s in audio_by_stem] + if not ctx: + return audio_by_stem + ctx_rms = float(np.sqrt(np.mean(np.square(np.sum(ctx, axis=0)))) + 1e-9) + for s, v in dec.items(): + if v == "generated" and s in audio_by_stem and s in GEN_ENERGY_RATIO: + a = audio_by_stem[s] + r = float(np.sqrt(np.mean(np.square(a))) + 1e-9) + audio_by_stem[s] = a * (GEN_ENERGY_RATIO[s] * ctx_rms / r) + return audio_by_stem + + +@torch.inference_mode() +def decode(latent): + return E._decode(models()["sa"], _fitT(latent)) # [2,S] cpu + +@torch.inference_mode() +def encode(audio_2S): + return encode_batch([audio_2S])[0] + + +@torch.inference_mode() +def encode_batch(auds): + """Encode many [2,S] clips in ONE SAME forward pass (CPU: far less overhead than N calls).""" + m = models(); xs = [] + for a in auds: + x = torch.as_tensor(a).float() + if x.shape[0] != 2: x = (x[:2] if x.shape[0] > 2 else x.repeat(2, 1)) + xs.append(x) + L = max(x.shape[1] for x in xs) + X = torch.stack([torch.nn.functional.pad(x, (0, L - x.shape[1])) for x in xs]).to(DEVICE) + Z = m["sa"].encode_audio(X, chunked=False) + return [Z[i].float().cpu() for i in range(Z.shape[0])] + + +@torch.inference_mode() +def decode_batch(lats): + """Decode many [256,T] latents in ONE SAME forward pass.""" + if not lats: + return [] + m = models() + Z = torch.stack([_fitT(l) for l in lats]).to(DEVICE).float() + A = m["sa"].decode_audio(Z) + return [A[i].clamp(-1, 1).cpu() for i in range(A.shape[0])] + + +def process_chunk(deg_latents: dict, mix_latent, present: dict): + """deg_latents/present keyed by stem. Restore present, generate missing. Returns + (final_latents dict, decisions dict).""" + final, dec = {}, {} + for s in STEMS: + if present.get(s) and s in deg_latents: + final[s] = restore(deg_latents[s], s, mix_latent); dec[s] = "restored" + ctx = [final[o] for o in final] # restored context + for s in STEMS: + if s in final: + continue + if s in GEN_TARGETS: # always invent missing bass + g = generate(ctx, s) if ctx else None + if g is not None: final[s] = g; dec[s] = "generated" + else: dec[s] = "absent (no generator/context)" + else: # drums/other: don't fabricate + dec[s] = "absent (kept out)" + return final, dec + + +# ---------------- viz / eval ---------------- +def spec_fig(audio_2S, title): + y = np.asarray(audio_2S).mean(0) if np.asarray(audio_2S).ndim == 2 else np.asarray(audio_2S) + S = librosa.amplitude_to_db(np.abs(librosa.stft(y, n_fft=1024, hop_length=256)) + 1e-6, ref=np.max) + fig, ax = plt.subplots(figsize=(4.2, 2.4)) + librosa.display.specshow(S, sr=SR, hop_length=256, x_axis="time", y_axis="log", ax=ax, cmap="magma") + ax.set_title(title, fontsize=9); ax.tick_params(labelsize=7) + fig.tight_layout() + return fig + +def multi_stft_np(a, b): + return E.multi_stft(torch.as_tensor(a).float(), torch.as_tensor(b).float()) + + +def load_audio_any(path): + """Read wav/flac via soundfile; fall back to librosa+ffmpeg for mp3/m4a/ogg/etc. + Returns ([C, S] float32, sr).""" + try: + a, sr = sf.read(path, dtype="float32") + return (a.T if a.ndim == 2 else a[None]), sr + except Exception: # mp3 & friends -> ffmpeg/audioread + y, sr = librosa.load(path, sr=None, mono=False) + y = np.asarray(y, dtype="float32") + return (y if y.ndim == 2 else y[None]), sr + + +# ---------------- val index ---------------- +def build_val_index(): + rows = {} + meta = VAL_CACHE / "metadata.csv" + for r in csv.DictReader(open(meta)): + rows.setdefault(r["pair_id"], {})[r["stem"]] = r + out = [] + for pid, st in rows.items(): + if not (SRC / "val/degraded" / f"{pid}.wav").exists(): + continue + present = [s for s in STEMS if st.get(s, {}).get("deg_present") == "1"] + # generation showcase = a MISSING bass that exists in clean (policy: always invent bass) + missing = [s for s in GEN_TARGETS if s not in present and st.get(s, {}).get("clean_present") == "1"] + if present and missing: # restore present + generate bass + out.append((pid, present, missing)) + out.sort(key=lambda x: (-len(x[2]), x[0])) + return out[:N_VAL] + +_VAL = None +def val_list(): + global _VAL + if _VAL is None: _VAL = build_val_index() + return _VAL + +def val_label(item): + pid, pres, miss = item + return f"{pid} | restore: {','.join(pres) or '-'} | generate: {','.join(miss) or '-'}" + + +# ---------------- pipelines ---------------- +def _empty_val(msg): + main = [None, None, None, None, None, None, msg] + stems = [] + for s in STEMS: + stems += [f"**{STEM_EMOJI[s]} {s}**", None, None] + return main + stems + + +def run_val(label): + if not label: + return _empty_val("Pick a sample and press Run.") + pid = label.split()[0] + deg_a, _ = sf.read(SRC / "val/degraded" / f"{pid}.wav", dtype="float32") + deg_a = deg_a.T if deg_a.ndim == 2 else deg_a[None] + clean_a, _ = sf.read(SRC / "val/clean" / f"{pid}.wav", dtype="float32") + clean_a = clean_a.T if clean_a.ndim == 2 else clean_a[None] + latdir = VAL_CACHE / pid / "latents" + deg_latents, present = {}, {} + for s in STEMS: + p = latdir / "degraded" / f"{s}.pt" + if p.exists(): + deg_latents[s] = torch.load(p, map_location="cpu").float(); present[s] = True + mix = torch.load(latdir / "degraded_mix.pt", map_location="cpu").float() + final, dec = process_chunk(deg_latents, mix, present) + fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys]) # one batched decode + after = {fkeys[i]: fdec[i].numpy() for i in range(len(fkeys))} + energy_match(after, dec) # level generated bass + bkeys = list(deg_latents); bdec = decode_batch([deg_latents[s] for s in bkeys]) + before = {bkeys[i]: bdec[i].numpy() for i in range(len(bkeys))} + out = sum(after.values()) # sum after leveling + n = min(out.shape[1], clean_a.shape[1], deg_a.shape[1]) + e_in = multi_stft_np(deg_a[:, :n].mean(0), clean_a[:, :n].mean(0)) + e_out = multi_stft_np(out[:, :n].mean(0), clean_a[:, :n].mean(0)) + n_rest = sum(v == "restored" for v in dec.values()); n_gen = sum(v == "generated" for v in dec.values()) + report = (f"**`{pid}`** ยท ๐Ÿ”ง restored {n_rest} ยท โœจ generated {n_gen} | " + f"full-mix vs clean (lower better): degraded **{e_in:.3f}** โ†’ output **{e_out:.3f}** " + f"{'โœ…' if e_out < e_in else 'โš ๏ธ'}") + main = [(SR, deg_a.T), (SR, out.T), (SR, clean_a.T), + spec_fig(deg_a, "degraded input"), spec_fig(out, "output"), spec_fig(clean_a, "clean reference"), + report] + stems = [] + for s in STEMS: + tag = {"restored": "๐Ÿ”ง restored", "generated": "โœจ generated"}.get(dec.get(s), "โ€”") + bef = (SR, before[s].T) if s in before else None + aft = (SR, after[s].T) if s in after else None + stems += [f"**{STEM_EMOJI[s]} {s}** โ€” {tag}", bef, aft] + return main + stems + + +def _empty_upload(msg): + main = [None, None, None, None, msg] + stems = [] + for s in STEMS: + stems += [f"**{STEM_EMOJI[s]} {s}**", None] + return main + stems + + +def _rms(x): + return float(np.sqrt(np.mean(x.astype(np.float64) ** 2)) + 1e-12) + + +def run_upload(filepath, start_s=0.0, end_s=0.0): + if not filepath: + return _empty_upload("Upload a wav/mp3 and press Run.") + a, sr = load_audio_any(filepath) # wav/flac/mp3/m4a... + if a.shape[0] == 1: a = np.repeat(a, 2, 0) + if sr != SR: + a = np.stack([librosa.resample(a[c], orig_sr=sr, target_sr=SR) for c in range(a.shape[0])]) + # explicit trim (seconds): actually chop the file before processing + dur = a.shape[1] / SR + s = max(0.0, float(start_s or 0.0)) + e = float(end_s) if (end_s and float(end_s) > 0) else dur + e = min(e, dur) + if e <= s: + s, e = 0.0, dur + a = a[:, int(s * SR):int(e * SR)] + trim_note = f"trim {s:.1f}โ€“{e:.1f}s of {dur:.1f}s ยท " + chunk = SR * 3; fade = SR // 2; hop = chunk - fade; total = a.shape[1] # 3s windows, 0.5s crossfade + try: + sep = get_separator() + except RuntimeError as e: # separation unavailable on this deployment + return _empty_upload(f"โš ๏ธ {e}") + + # ---- pass 1: process OVERLAPPING 3s windows (hop 2.5s) independently ---- + windows = []; summ = {} + for i in range(0, max(1, total), hop): + seg = a[:, i:i + chunk]; orig = seg.shape[1] + if orig <= 0: + break + if orig < chunk: + seg = np.pad(seg, ((0, 0), (0, chunk - orig))) # pad silence then crop output back + stems = separate_chunk(sep, seg) + present_stems = [s for s, au in stems.items() + if (lambda rp: rp[0] > PRESENT_RMS or rp[1] > PRESENT_PEAK)(rms_peak_db(au.mean(0)))] + enc = encode_batch([stems[s] for s in present_stems] + [seg]) # one batched encode (stems + mix) + deg_latents = {present_stems[j]: enc[j] for j in range(len(present_stems))} + present = {s: True for s in present_stems}; mix_lat = enc[-1] + final, dec = process_chunk(deg_latents, mix_lat, present) + for s, v in dec.items(): summ[s] = v + fkeys = list(final); fdec = decode_batch([final[s] for s in fkeys]) # one batched decode + chunk_audio = {s: fdec[j].numpy() for j, s in enumerate(fkeys)} # SAME decode = fixed length/window + energy_match(chunk_audio, dec) # intra-chunk: level generated bass + windows.append({"start": i, "orig": orig, "stems": chunk_audio}) + if i + chunk >= total: + break + + # ---- GLUE 1 (profile): gently match each window's per-stem level to the song-wide MEDIAN, so chunks + # share a consistent profile instead of drifting. alpha<1 preserves real dynamics; clip = no artifacts. + ALPHA, LO, HI = 0.5, 0.7, 1.4 + for s in STEMS: + rmss = [_rms(w["stems"][s]) for w in windows if s in w["stems"] and _rms(w["stems"][s]) > 1e-5] + if len(rmss) >= 2: + tgt = float(np.median(rmss)) + for w in windows: + if s in w["stems"] and _rms(w["stems"][s]) > 1e-5: + w["stems"][s] = w["stems"][s] * float(np.clip((tgt / _rms(w["stems"][s])) ** ALPHA, LO, HI)) + + # ---- GLUE 2 (seams): crossfaded overlap-add, done in OUTPUT-sample space (SAME decodes each 3s window + # to a fixed length Ld, not the input `orig`). out_hop/out_fade scale the input hop/fade by Ld/chunk. + # One global weight timeline (incl. windows where a stem is absent) so a stem fades across a boundary + # into a neighbour that lacks it. ---- + n = len(windows) + Ld = next((au.shape[1] for w in windows for au in w["stems"].values()), 0) + if Ld == 0: + return _empty_upload(f"{trim_note}no restorable content found.") + out_fade = min(int(round(fade * Ld / chunk)), Ld // 2) + out_hop = max(1, Ld - out_fade) + def _Lout(w): # real output length (last/short window cropped) + return Ld if w["orig"] >= chunk else min(Ld, int(round(w["orig"] * Ld / chunk))) + starts = [k * out_hop for k in range(n)] + total_out = starts[-1] + _Lout(windows[-1]) + Wg = np.zeros(total_out); envs = [] + for k, w in enumerate(windows): + L = _Lout(w); env = np.ones(L) + if k > 0 and out_fade > 0: env[:min(out_fade, L)] = np.linspace(0, 1, out_fade, endpoint=False)[:min(out_fade, L)] + if k < n - 1 and L >= out_fade and out_fade > 0: env[-out_fade:] = np.linspace(1, 0, out_fade, endpoint=False) + envs.append(env); Wg[starts[k]:starts[k] + L] += env + Wg = np.maximum(Wg, 1e-6) + per_stem_audio = {} + for s in STEMS: + if not any(s in w["stems"] for w in windows): + continue + buf = np.zeros((2, total_out)) + for k, w in enumerate(windows): + if s in w["stems"]: + L = _Lout(w); buf[:, starts[k]:starts[k] + L] += w["stems"][s][:, :L] * envs[k] + per_stem_audio[s] = (buf / Wg).astype(np.float32) + out = sum(per_stem_audio.values()) if per_stem_audio else np.zeros((2, total_out), np.float32) + + rep = (f"{trim_note}**{n}** overlapping 3 s windows ยท hop {hop/SR:.1f}s, {fade/SR:.1f}s crossfade ยท " + "glued to song profile. " + " ยท ".join(f"{STEM_EMOJI[st]} {summ.get(st, '-')}" for st in STEMS)) + main = [(SR, a.T), (SR, out.T), spec_fig(a, "input"), spec_fig(out, "output"), rep] + stems_out = [] + for s in STEMS: + tag = {"restored": "๐Ÿ”ง restored", "generated": "โœจ generated"}.get(summ.get(s), "โ€”") + au = (SR, per_stem_audio[s].T) if s in per_stem_audio else None + stems_out += [f"**{STEM_EMOJI[s]} {s}** โ€” {tag}", au] + return main + stems_out + + +# ---------------- separator (uploads only) ---------------- +# EXACT-PARITY: the train/val cache was built with htdemucs v4 (audio-separator's htdemucs.yaml, +# shifts=1 overlap=0.25 โ€” see demucs.py). Live upload MUST use those SAME v4 weights or the restorer +# sees a different stem distribution. Priority: (1) audio-separator (exactly how the cache was built โ€” +# present locally); (2) the official `demucs` package (SAME htdemucs v4 weights, far lighter deps so it +# installs on the CPU Space without the numpy/pywt clash that disabled this before). We deliberately do +# NOT fall back to torchaudio's HDemucs โ€” that's v3 (no transformer), i.e. NOT identical stems. +_SEP = {} +def get_separator(): + if _SEP: return _SEP + try: # (1) audio-separator โ€” exact, as used to build the cache + from audio_separator.separator import Separator + s = Separator(output_dir="/tmp/app_sep", model_file_dir="/tmp/audio-separator-models") + s.load_model("htdemucs.yaml") + _SEP.update(kind="audio_separator", sep=s) + return _SEP + except Exception: + pass + try: # (2) demucs package โ€” SAME htdemucs v4 weights + from demucs.pretrained import get_model # NB: on the Space there is no root demucs.py to shadow this + from demucs.apply import apply_model + m = get_model("htdemucs").to(DEVICE).eval() + _SEP.update(kind="demucs", model=m, apply=apply_model, sources=list(m.sources)) + return _SEP + except Exception as e: + raise RuntimeError( + "Live upload separation is unavailable: htdemucs v4 (audio-separator or the demucs package) " + "isn't installed here. Use the ๐ŸŽง Val samples tab โ€” full restore+generate pipeline. " + "(torchaudio's HDemucs is v3 and would NOT match the trained stems.)" + ) from e + +def separate_chunk(sep, seg_2S): + """seg_2S [2,S] @ SR -> {stem: [2,S] float32}, via the same htdemucs v4 weights the cache used.""" + if sep["kind"] == "audio_separator": + tmp = "/tmp/app_sep_in.wav"; sf.write(tmp, seg_2S.T, SR) + files = sep["sep"].separate(tmp) + out = {} + for f in files: + fl = f.lower() + for s in STEMS: + if s in fl: + au, _ = sf.read(os.path.join("/tmp/app_sep", f) if not os.path.isabs(f) else f, dtype="float32") + out[s] = (au.T if au.ndim == 2 else np.repeat(au[None], 2, 0)) + return out + # demucs package: per-mix normalize (the standard htdemucs recipe), shifts=1 overlap=0.25 to match. + wav = torch.as_tensor(seg_2S, dtype=torch.float32, device=DEVICE) # [2,S] + ref = wav.mean(0); mu, sd = ref.mean(), ref.std().clamp_min(1e-8) + with torch.inference_mode(): + src = sep["apply"](sep["model"], ((wav - mu) / sd)[None], + shifts=1, overlap=0.25, split=True, device=DEVICE)[0] # [4,2,S] + src = (src * sd + mu).cpu().numpy() + return {s: src[i] for i, s in enumerate(sep["sources"]) if s in STEMS} + + +EXPLAIN = """ +## How this works (high level) + +**The problem.** Old/lo-fi recordings are *degraded* (muffled, missing instruments). We work in +**SAME-L's latent space** โ€” a pretrained **neural audio autoencoder** (a learned codec, ร  la the +VAEs behind latent diffusion) that compresses 3 s of stereo into a small `256ร—32` "summary" +(~11 frames/sec). Everything below operates on these latents, then decodes back to audio โ€” so the +models stay tiny and fast. + +**The models, in plain terms:** +- **SAME-L** โ€” the compact "latent" space everything runs in (a frozen neural audio codec: encode โ†” decode). +- **htdemucs** โ€” splits a song into its 4 instrument stems. +- **Restorer** โ€” *cleans up* a muffled/damaged stem so it sounds clear again. It fixes what's there; it doesn't add new parts. (A deterministic Conformer-style net, sharpened by the GAN below.) +- **Generator** โ€” *invents* a missing instrument (mainly bass) that fits the song, like an AI session musician. (Conditional flow-matching.) +- **GAN "judge"** โ€” rates how *real* the final mix sounds; the restorer and generator compete to fool it, which pushes them toward genuinely real audio instead of a safe, dull average. + +**Step 1 โ€” Separate.** htdemucs splits the input mix into 4 stems: *other, vocals, drums, bass*. + +**Step 2 โ€” Route.** For each stem we ask: *is it actually there?* (energy gate). +- **Present** โ†’ it just needs cleanup โ†’ **RESTORE**. +- **Missing bass** โ†’ every song has a low end, so an absent bass is **invented** โ†’ **GENERATE**. +- **Missing drums/other** โ†’ a song may legitimately have none, so we **do not fabricate** them โ€” restore-only. + +**Step 3a โ€” Restorer (deterministic).** For present stems, a small **Conformer-style network** +(local convolutions for texture **+ global self-attention** so each moment sees the whole clip) +nudges the degraded latent toward a clean one, as a *residual* on the input. It's *deterministic* +because the degraded signal already tells us the answer โ€” we just sharpen it. Trained with a +**stem-balanced** loss so quiet stems (bass) aren't drowned out by loud ones. + +**Step 3b โ€” Generator (creative, bass).** When the bass is missing there is no answer to copy โ€” +many different bass lines could fit. So this is a **conditional flow-matching** generator: starting +from noise it follows a learned velocity field to *sample* a plausible **bass conditioned on the +other (restored) stems**, so it locks to the song's harmony. **Classifier-free guidance** lets us +dial how strongly it commits to the accompaniment. (Probes showed bass is strongly determined by the +accompaniment and the codec preserves it well; drums are timing-limited by the ~11 Hz latent, so we +don't fabricate drums.) + +**Step 4 โ€” Decode & remix.** Each final latent is decoded to audio and summed into the output mix. + +> Restoration leans on the *existing* signal (works best when degradation is mild). Generation +> *creates* absent instruments to accompany what's there. Inference only ever sees restored stems โ€” +> never the clean originals. +""" + + +THEME = gr.themes.Base( + primary_hue=gr.themes.colors.orange, + secondary_hue=gr.themes.colors.orange, + neutral_hue=gr.themes.colors.zinc, + font=[gr.themes.GoogleFont("Inter"), "system-ui", "sans-serif"], + radius_size=gr.themes.sizes.radius_sm, +).set( + # dark look as the default (set the light variants to dark so it's black/orange regardless of OS) + body_background_fill="#0a0a0b", + body_text_color="#e9e9ec", + body_text_color_subdued="#8a8a93", + background_fill_primary="#151518", + background_fill_secondary="#1c1c20", + block_background_fill="#151518", + block_border_color="#272730", + block_label_background_fill="#151518", + block_label_text_color="#ff8a3d", + block_title_text_color="#f2f2f4", + border_color_primary="#272730", + input_background_fill="#1c1c20", + button_primary_background_fill="#ff7a18", + button_primary_background_fill_hover="#ff9344", + button_primary_text_color="#0a0a0b", + button_secondary_background_fill="#272730", + button_secondary_text_color="#e9e9ec", + color_accent_soft="#2a1a0e", + slider_color="#ff7a18", +) + +CSS = """ +.gradio-container {max-width: 1080px !important; margin: auto !important; background:#0a0a0b;} +.fullmix audio {width: 100%;} +h1,h2,h3,h4 {color:#f4f4f6 !important; letter-spacing:.2px;} +h1 {font-weight:700; border-left:4px solid #ff7a18; padding-left:12px;} +h4 {color:#ff9a55 !important; text-transform:uppercase; font-size:.8rem; letter-spacing:.6px;} +a {color:#ff8a3d !important;} +footer {display:none !important;} +::-webkit-scrollbar{width:9px;height:9px} ::-webkit-scrollbar-thumb{background:#3a3a44;border-radius:6px} +""" + + +def _mix_column(title, with_plot=True): + gr.Markdown(f"#### {title}") + a = gr.Audio(label=None, show_label=False, elem_classes="fullmix") + p = gr.Plot(show_label=False) if with_plot else None + return a, p + + +def build_ui(): + # NOTE: do NOT load models here โ€” that delays the port bind ~15 s and makes the + # auto-opened browser tab hit a dead port. UI uses cheap dir listing; models load + # lazily on first action (warmed in a background thread at launch). + rests = list_runs("restorer"); gens = list_runs("generator") + dr, dg = _default_runs(rests, gens) + with gr.Blocks(title="SAME Restore + Generate", theme=THEME, css=CSS) as demo: + gr.Markdown("# ๐ŸŽ›๏ธ Stem Restoration + Generation\n" + "Upload a muffled/lo-fi song. We split it into instruments, **clean up** the ones that are " + "there, **invent** a missing bass that fits, and remix โ€” all in a compact neural-audio space. " + "A GAN 'judge' keeps the result sounding *real*, not dull.") + status = gr.Markdown(status_md()) + with gr.Accordion("โš™๏ธ Choose models", open=False): + with gr.Row(): + rest_dd = gr.Dropdown(rests, value=dr, label="Restorer", scale=2) + gen_dd = gr.Dropdown(gens + ["(none)"], value=dg, label="Generator (bass)", scale=2) + load_btn = gr.Button("โ†ป Load", variant="primary", scale=1) + load_btn.click(set_models, [rest_dd, gen_dd], status) + with gr.Tab("๐ŸŽง Val samples"): + with gr.Row(): + dd = gr.Dropdown([val_label(x) for x in val_list()], label="3 s val sample (each shows both restore + generate)", scale=5) + btn = gr.Button("โ–ถ Run", variant="primary", scale=1) + rep = gr.Markdown() + with gr.Row(equal_height=True): + with gr.Column(): a_in, s_in = _mix_column("๐ŸŽš๏ธ Degraded input") + with gr.Column(): a_out, s_out = _mix_column("โœจ Output") + with gr.Column(): a_ref, s_ref = _mix_column("๐ŸŽฏ Clean reference") + with gr.Accordion("๐Ÿ”ฌ Per-stem detail (before โ†’ after)", open=False): + comps = [] + for s in STEMS: + with gr.Row(equal_height=True): + lbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**") + bef = gr.Audio(label="before (degraded)", show_label=True, scale=2) + aft = gr.Audio(label="after (restored/generated)", show_label=True, scale=2) + comps += [lbl, bef, aft] + btn.click(run_val, dd, [a_in, a_out, a_ref, s_in, s_out, s_ref, rep] + comps) + with gr.Tab("โฌ†๏ธ Upload"): + with gr.Row(): + up = gr.Audio(label="Upload wav/mp3 (chunked into 3 s; last chunk padded then cropped)", + type="filepath", sources=["upload"], editable=True, + waveform_options=gr.WaveformOptions( + waveform_color="#5a5a66", waveform_progress_color="#ff7a18", + trim_region_color="#ff7a18"), + scale=5) + ub = gr.Button("โ–ถ Run", variant="primary", scale=1) + with gr.Row(): + trim_start = gr.Number(value=0, label="Trim start (s)", scale=1) + trim_end = gr.Number(value=0, label="Trim end (s, 0 = to end)", scale=1) + gr.Markdown("*Set start/end to actually chop the file before processing " + "(leave both 0 to process the whole upload).*") + urep = gr.Markdown() + with gr.Row(equal_height=True): + with gr.Column(): ua_in, us_in = _mix_column("โฌ†๏ธ Input") + with gr.Column(): ua_out, us_out = _mix_column("โœจ Output") + with gr.Accordion("๐Ÿ”ฌ Per-stem detail (output)", open=False): + ucomps = [] + for s in STEMS: + with gr.Row(equal_height=True): + ulbl = gr.Markdown(f"**{STEM_EMOJI[s]} {s}**") + uaft = gr.Audio(label="output stem", show_label=True, scale=3) + ucomps += [ulbl, uaft] + ub.click(run_upload, [up, trim_start, trim_end], [ua_in, ua_out, us_in, us_out, urep] + ucomps) + with gr.Accordion("โ„น๏ธ How this works (restorer + generator, intuitively)", open=False): + gr.Markdown(EXPLAIN) + return demo + + +def _free_port(port): + """Kill any process still holding the port so re-launch binds cleanly.""" + try: + out = os.popen(f"fuser {port}/tcp 2>/dev/null").read().split() + for pid in out: + if pid.strip().isdigit() and int(pid) != os.getpid(): + os.system(f"kill -9 {pid} 2>/dev/null") + except Exception: + pass + + +if __name__ == "__main__": + port = int(os.environ.get("PORT", 7860)) + _free_port(port) # clear a stale bind from a previous run + demo = build_ui() + import threading + threading.Thread(target=models, daemon=True).start() # warm models in bg; port still binds instantly + try: + demo.launch(server_name="0.0.0.0", server_port=port, share=False, + ssr_mode=False) # ssr_mode off: Gradio-6 SSR hangs without a node runtime + except KeyboardInterrupt: + print("\n[app] shutting downโ€ฆ") + finally: # flush the port on exit (Ctrl+C / close) + try: demo.close() + except Exception: pass + try: gr.close_all() + except Exception: pass + _free_port(port) + print("[app] port released.") diff --git a/restoflow/config.py b/restoflow/config.py new file mode 100644 index 0000000000000000000000000000000000000000..4766c63ddd19b06b1227afe4f4041d7ccb0293db --- /dev/null +++ b/restoflow/config.py @@ -0,0 +1,70 @@ +"""Central config for the SAME-latent restoration model. + +Validated facts driving these defaults (see project memory / restore_demo): + - degraded->clean is one-to-many => generative (flow-matching), not regression. + - cosine distance in the 256-latent space tracks decoded-audio distance + (within-clip Spearman ~0.94) => cheap no-decode train/val metric. + - one shared model conditioned on stem id beats per-stem-separate nets (helps + low-data drums/bass). So: presence ROUTER + ONE conditioned flow. + - drums are legitimately low-RMS (transient) => presence = rms OR peak. +""" +from dataclasses import dataclass, field +from pathlib import Path + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") +STEMS = ("other", "vocals", "drums", "bass") +STEM_ID = {s: i for i, s in enumerate(STEMS)} + + +@dataclass +class Cfg: + # ---- data ---- + cache_roots: tuple = (str(BASE / "demucs_results_all"),) # train cache; override with --cache-roots + val_cache_roots: tuple = (str(BASE / "demucs_results_val_full"),) # held-out val (disjoint source split) + out_dir: str = str(BASE / "restoflow_runs" / "proto") + stems: tuple = STEMS + T: int = 32 # fixed frames per 3s clip (crop/pad) + latent_dim: int = 256 + val_frac: float = 0.1 + split_seed: int = 0 + # ---- hardened presence (training filter; recomputed from metadata columns) ---- + rms_thr: float = -40.0 + peak_thr: float = -25.0 + max_drop_db: float = 15.0 # degraded stem >this dB below clean = gutted -> drop (generation) + # ---- normalization (per-stem per-dim z-score from TRAIN clean) ---- + stats_path: str = "" # default: /norm_stats.pt + # ---- model (small; latents are tiny) ---- + model_kind: str = "det" # "det" = deterministic mix+stem restorer (validated best); + # "attn" = det + global Conformer-lite attention (experiment); + # "flow" = conditional flow-matching (kept for research) + use_mix: bool = True # condition on degraded full-mix latent (big drums win) + hidden: int = 384 + depth: int = 5 + stem_emb_dim: int = 32 + heads: int = 4 # attn restorer (model_kind="attn") only + sigma: float = 0.0 # flow-only noisy-start std (sweep showed 0 best; noise hurts) + # ---- v4: stem balancing + energy matching (fight bass-drowning + dullness) ---- + # VALIDATED (v4b): alpha=0.5 is the sweet spot (un-drowns bass without starving loud + # stems); alpha=1.0 overshoots (worse on every stem). energy_w MUST stay 0: the + # magnitude penalty rewards loud-but-wrong outputs -> regresses below degraded (v4). + stem_weight_alpha: float = 0.5 # 0=uniform (v3), 0.5=best, 1=full inverse-freq (worse) + energy_w: float = 0.0 # KEEP 0 (validated harmful); dullness needs a stochastic fix, not this + # ---- distributional aux loss (cheap latent FAD-surrogate; fights deterministic dullness) ---- + dist_loss: str = "none" # none | var (diag mean+var) | mmd (RBF) | moment (full cov, noisy) + dist_loss_w: float = 0.0 # weight of the per-stem distributional term added to det/attn MSE + # ---- train ---- + batch: int = 128 + lr: float = 1e-3 + wd: float = 1e-4 + epochs: int = 200 + device: str = "cuda" + num_workers: int = 4 + # ---- eval / demo ---- + eval_every: int = 20 + sample_steps: int = 40 + eval_max_pairs: int = 64 + demo_pairs: int = 6 + model_id: str = "stabilityai/SAME-L" + + def stats_file(self) -> Path: + return Path(self.stats_path) if self.stats_path else Path(self.out_dir) / "norm_stats.pt" diff --git a/restoflow/data.py b/restoflow/data.py new file mode 100644 index 0000000000000000000000000000000000000000..70d6ec3c6262a50687b81ff194e49ab757240fd0 --- /dev/null +++ b/restoflow/data.py @@ -0,0 +1,136 @@ +"""Dataset + normalization stats for SAME-latent restoration. + +- Reads metadata.csv from one or more cache roots. +- Keeps only RESTORATION (pair,stem) under hardened presence (router.is_restoration). +- Deterministic by-pair train/val split (or explicit val_cache_roots). +- Per-stem per-dim z-score normalization from TRAIN clean latents. +Item: dict(deg[256,T], clean[256,T] standardized, stem_id, stem, pair_id, root). +""" +from __future__ import annotations +import csv +from pathlib import Path +import torch +from torch.utils.data import Dataset + +from .config import Cfg, STEM_ID +from . import router + + +def _index(cfg: Cfg, roots, split_mode): + """Collect item dicts. split_mode: 'train'/'val'/'all'. If cfg.val_cache_roots set, + roots are used as-is for the requested split; else by-pair hash split.""" + items = [] + explicit_val = bool(cfg.val_cache_roots) + for root in roots: + root = Path(root) + meta = root / "metadata.csv" + if not meta.exists(): + continue + for row in csv.DictReader(open(meta)): + if row.get("stem") not in cfg.stems: + continue + if not router.is_restoration(row, cfg.rms_thr, cfg.peak_thr, cfg.max_drop_db): + continue + if row.get("clean_latent_path") in ("", None) or row.get("deg_latent_path") in ("", None): + continue + if not explicit_val and split_mode in ("train", "val"): + if router.pair_split(row["pair_id"], cfg.val_frac, cfg.split_seed) != split_mode: + continue + mix = root / row["pair_id"] / "latents" / "degraded_mix.pt" + if cfg.use_mix and not mix.exists(): + continue + items.append({ + "pair_id": row["pair_id"], "stem": row["stem"], "root": str(root), + "clean": str(root / row["clean_latent_path"]), + "deg": str(root / row["deg_latent_path"]), + "mix": str(mix), + }) + return items + + +def _fit_T(x: torch.Tensor, T: int) -> torch.Tensor: # x [256, t] -> [256, T] + t = x.shape[1] + if t == T: + return x + if t > T: + return x[:, :T] + return torch.cat([x, x[:, -1:].repeat(1, T - t)], dim=1) # edge-repeat pad + + +class LatentRestore(Dataset): + def __init__(self, cfg: Cfg, split: str, stats: dict): + self.cfg = cfg + roots = cfg.val_cache_roots if (split == "val" and cfg.val_cache_roots) else cfg.cache_roots + self.items = _index(cfg, roots, split) + self.stats = stats # stem -> {'mu':[256], 'sd':[256]} + + def __len__(self): + return len(self.items) + + def _norm(self, x, stem): + mu, sd = self.stats[stem]["mu"], self.stats[stem]["sd"] + return (x - mu[:, None]) / sd[:, None] + + def __getitem__(self, i): + it = self.items[i] + clean = _fit_T(torch.load(it["clean"], map_location="cpu").float(), self.cfg.T) + deg = _fit_T(torch.load(it["deg"], map_location="cpu").float(), self.cfg.T) + out = { + "deg": self._norm(deg, it["stem"]), + "clean": self._norm(clean, it["stem"]), + "stem_id": torch.tensor(STEM_ID[it["stem"]], dtype=torch.long), + "stem": it["stem"], "pair_id": it["pair_id"], + } + if self.cfg.use_mix: + mix = _fit_T(torch.load(it["mix"], map_location="cpu").float(), self.cfg.T) + m = self.stats["__mix__"] + out["mix"] = (mix - m["mu"][:, None]) / m["sd"][:, None] + return out + + +def build_norm_stats(cfg: Cfg) -> dict: + """Per-stem per-dim mean/std from TRAIN clean latents. Saved to cfg.stats_file().""" + items = _index(cfg, cfg.cache_roots, "train") + acc = {s: {"n": 0, "sum": torch.zeros(cfg.latent_dim, dtype=torch.float64), + "sq": torch.zeros(cfg.latent_dim, dtype=torch.float64)} for s in cfg.stems} + for it in items: + x = torch.load(it["clean"], map_location="cpu").double() # [256,t] + a = acc[it["stem"]] + a["n"] += x.shape[1]; a["sum"] += x.sum(1); a["sq"] += (x * x).sum(1) + stats = {} + for s, a in acc.items(): + if a["n"] == 0: + stats[s] = {"mu": torch.zeros(cfg.latent_dim), "sd": torch.ones(cfg.latent_dim), "n": 0} + continue + mu = a["sum"] / a["n"] + var = (a["sq"] / a["n"] - mu * mu).clamp_min(1e-8) + stats[s] = {"mu": mu.float(), "sd": var.sqrt().float(), "n": a["n"]} + if cfg.use_mix: # stem-independent mix stats (one degraded_mix per pair) + n = 0; ssum = torch.zeros(cfg.latent_dim, dtype=torch.float64); ssq = torch.zeros(cfg.latent_dim, dtype=torch.float64) + seen = set() + for it in items: + if it["pair_id"] in seen or not Path(it["mix"]).exists(): + continue + seen.add(it["pair_id"]); x = torch.load(it["mix"], map_location="cpu").double() + n += x.shape[1]; ssum += x.sum(1); ssq += (x * x).sum(1) + if n: + mu = ssum / n; var = (ssq / n - mu * mu).clamp_min(1e-8) + stats["__mix__"] = {"mu": mu.float(), "sd": var.sqrt().float(), "n": n} + out = cfg.stats_file(); out.parent.mkdir(parents=True, exist_ok=True) + torch.save(stats, out) + print(f"[stats] frames/stem:", {s: stats[s]["n"] for s in cfg.stems}, "-> saved", out) + return stats + + +def load_or_build_stats(cfg: Cfg) -> dict: + f = cfg.stats_file() + return torch.load(f) if f.exists() else build_norm_stats(cfg) + + +if __name__ == "__main__": + c = Cfg() + build_norm_stats(c) + for sp in ("train", "val"): + ds = LatentRestore(c, sp, load_or_build_stats(c)) + from collections import Counter + print(sp, "items:", len(ds), dict(Counter(it["stem"] for it in ds.items))) diff --git a/restoflow/distloss.py b/restoflow/distloss.py new file mode 100644 index 0000000000000000000000000000000000000000..10d3cd71082fe067c6537857f2cd85c98d809dad --- /dev/null +++ b/restoflow/distloss.py @@ -0,0 +1,69 @@ +"""Cheap DISTRIBUTIONAL losses in SAME-latent space โ€” "the latent is the mirror" of the audio, +so matching latent distributions trains the model to produce the DISTRIBUTION of plausible clean +stems (the one-to-many task), not a single conditional-mean target (which oversmooths/dulls). + +Differentiable surrogates for FAD (FAD is eval-only: its matrix-sqrt is unstable as a loss): + - moment_match_loss : ||mu_p - mu_c||^2 + ||Cov_p - Cov_c||_F^2 = Frechet WITHOUT the sqrt. + Cheap (mean + 256x256 cov), stable. (cf. Gram/style loss; CMD.) + - mmd_loss : multi-bandwidth RBF MMD^2 (Gretton); GMMN-style moment matching. +Each latent FRAME is a sample: [B,256,T] -> [B*T,256]. Use as an AUX term added to MSE/flow loss +(MSE keeps per-sample alignment; this term restores the spread/detail the mean collapses). +""" +from __future__ import annotations +import torch + + +def _flat(x: torch.Tensor) -> torch.Tensor: # [B,D,T] -> [B*T, D] + return x.permute(0, 2, 1).reshape(-1, x.shape[1]) + + +def var_match_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + """DIAGONAL moment match: ||mu_p-mu_c||^2 + ||var_p-var_c||^2 (per latent dim). + Low-variance estimator (256 stats, not 256x256) -> usable at batch scale. Directly targets + DULLNESS = the deterministic restorer's per-dim variance deficit (mean-collapse).""" + P, C = _flat(pred), _flat(target) + dmu = (P.mean(0) - C.mean(0)).pow(2).sum() + dvar = (P.var(0, unbiased=False) - C.var(0, unbiased=False)).pow(2).sum() + return dmu + dvar + + +def moment_match_loss(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + """||mu_p-mu_c||^2 + ||Cov_p-Cov_c||_F^2 (full FAD surrogate, no matrix sqrt). + NOTE: the full 256x256 cov is noise-dominated at batch scale โ€” prefer var_match/mmd.""" + P, C = _flat(pred), _flat(target) + dmu = (P.mean(0) - C.mean(0)).pow(2).sum() + Pc, Cc = P - P.mean(0, keepdim=True), C - C.mean(0, keepdim=True) + covP = (Pc.t() @ Pc) / (P.shape[0] - 1) + covC = (Cc.t() @ Cc) / (C.shape[0] - 1) + return dmu + (covP - covC).pow(2).sum() + + +def mmd_loss(pred: torch.Tensor, target: torch.Tensor, bandwidths=(2., 4., 8., 16., 32.), + max_n: int = 1024) -> torch.Tensor: + """Unbiased-ish multi-bandwidth RBF MMD^2 between latent-frame sets. Subsamples to max_n + frames per side to keep the O(N^2) kernel cheap.""" + P, C = _flat(pred), _flat(target) + if P.shape[0] > max_n: + P = P[torch.randperm(P.shape[0], device=P.device)[:max_n]] + if C.shape[0] > max_n: + C = C[torch.randperm(C.shape[0], device=C.device)[:max_n]] + + def k(a, b): + d2 = torch.cdist(a, b).pow(2) + return sum(torch.exp(-d2 / (2 * h * h)) for h in bandwidths) + + return k(P, P).mean() + k(C, C).mean() - 2 * k(P, C).mean() + + +def dist_term(kind: str, pred: torch.Tensor, target: torch.Tensor, sid: torch.Tensor) -> torch.Tensor: + """Per-stem distributional aux loss (matches each stem's distribution separately).""" + fn = {"var": var_match_loss, "moment": moment_match_loss, "mmd": mmd_loss}.get(kind) + if fn is None: + return pred.new_zeros(()) + tot, n = pred.new_zeros(()), 0 + for s in sid.unique(): + m = sid == s + if m.sum() < 2: + continue + tot = tot + fn(pred[m], target[m]); n += 1 + return tot / max(n, 1) diff --git a/restoflow/eval.py b/restoflow/eval.py new file mode 100644 index 0000000000000000000000000000000000000000..c5ce11cf39c2221b27ca43ddac3f72e0c310fa30 --- /dev/null +++ b/restoflow/eval.py @@ -0,0 +1,171 @@ +"""Evaluation: per-stem latent + decoded-audio metrics, and REMIX (sum stems) vs true clean. + +Decodes through frozen SAME. Cheap latent metrics (cosine) are the primary signal; +decoded multi-STFT + remix are the honest audio-domain checks. Also dumps demo audio. +""" +from __future__ import annotations +from collections import defaultdict +from pathlib import Path +import warnings +import numpy as np +import torch +import soundfile as sf + +from .config import Cfg, STEM_ID +from . import data as D +from . import flow as F +from . import fad + + +def load_same(cfg: Cfg): + warnings.filterwarnings("ignore") + from stable_audio_tools import get_pretrained_model + sa, sacfg = get_pretrained_model(cfg.model_id) + sa = sa.to(cfg.device).eval() + for p in sa.parameters(): + p.requires_grad_(False) + return sa, int(sacfg.get("sample_rate", 44100)) + + +def _destd(z_std, stats, stem): # [256,T] standardized -> raw + mu, sd = stats[stem]["mu"], stats[stem]["sd"] + return z_std * sd[:, None].to(z_std.device) + mu[:, None].to(z_std.device) + + +@torch.inference_mode() +def _decode(sa, raw_latent): # [256,T] -> stereo [2,S] cpu + a = sa.decode_audio(raw_latent[None].to(next(sa.parameters()).device).float()) + return a[0].clamp(-1, 1).cpu() + + +def multi_stft(a, b): # mono tensors + tot = 0.0 + for nf in (512, 1024, 2048): + Sa = torch.stft(a, nf, hop_length=nf // 4, return_complex=True).abs() + Sb = torch.stft(b, nf, hop_length=nf // 4, return_complex=True).abs() + tot += (torch.log(Sa + 1e-5) - torch.log(Sb + 1e-5)).abs().mean().item() + return tot / 3 + + +def _select_pairs(by_pair, budget): + """Stratified pair selection. budget<=0 -> all pairs. Otherwise front-load pairs + that contain rare stems (bass/drums) so every stem gets eval coverage, not just + the head of the list. Order each pair by the rarity of its scarcest stem.""" + pairs = list(by_pair.items()) + if budget is None or budget <= 0 or budget >= len(pairs): + return pairs + freq = defaultdict(int) + for _, its in pairs: + for it in its: + freq[it["stem"]] += 1 + pairs.sort(key=lambda kv: min(freq[it["stem"]] for it in kv[1])) # rarest-stem pairs first + return pairs[:budget] + + +def evaluate(cfg: Cfg, model, stats, sa, sr, epoch=0): + model.eval() + items = D._index(cfg, cfg.val_cache_roots or cfg.cache_roots, "val") + by_pair = defaultdict(list) + for it in items: + by_pair[it["pair_id"]].append(it) + pairs = _select_pairs(by_pair, cfg.eval_max_pairs) + + per = defaultdict(lambda: defaultdict(list)) # stem -> metric -> [vals] + remix = defaultdict(list) + flat = {k: defaultdict(list) for k in ("clean", "in", "out")} # latent-FAD: stem -> [ [256,T] ] + frmx = {"clean": [], "in": [], "out": []} # latent-FAD REMIX (per-pair latent sum) + demo_root = Path(cfg.out_dir) / "demo" / f"epoch_{epoch:04d}" + + for pi, (pid, its) in enumerate(pairs): + clean_mix = deg_mix = rest_mix = None + clat = dlat = rlat = None + for it in its: + stem = it["stem"]; sid = torch.tensor([STEM_ID[stem]], device=cfg.device) + raw_c = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) + raw_d = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) + mu, sd = stats[stem]["mu"], stats[stem]["sd"] + std_d = ((raw_d - mu[:, None]) / sd[:, None]).to(cfg.device)[None] + std_c = ((raw_c - mu[:, None]) / sd[:, None]).to(cfg.device)[None] + std_mix = None + if cfg.use_mix: + raw_m = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) + m = stats["__mix__"] + std_mix = ((raw_m - m["mu"][:, None]) / m["sd"][:, None]).to(cfg.device)[None] + if cfg.model_kind in ("det", "attn"): + std_r = model(std_d, sid, std_mix) + elif cfg.model_kind == "flowattn": + std_r = F.sample_mix(model, std_d, std_mix, sid, cfg.sample_steps, cfg.sigma) + else: + std_r = F.sample(model, std_d, sid, cfg.sample_steps, cfg.sigma) + + per[stem]["cos_in"].append(F.latent_cos_dist(std_d, std_c).item()) # baseline (degraded) + per[stem]["cos_out"].append(F.latent_cos_dist(std_r, std_c).item()) # restored + per[stem]["relL2_out"].append(F.latent_relL2(std_r, std_c).item()) + per[stem]["energy_out"].append(F.energy_ratio(std_r, std_c).item()) + + raw_r = _destd(std_r[0], stats, stem) + rr = raw_r.detach().cpu() # latent-FAD (raw latents, no decode) + flat["clean"][stem].append(raw_c); flat["in"][stem].append(raw_d); flat["out"][stem].append(rr) + clat = raw_c if clat is None else clat + raw_c + dlat = raw_d if dlat is None else dlat + raw_d + rlat = rr if rlat is None else rlat + rr + ac = _decode(sa, raw_c); ad = _decode(sa, raw_d); ar = _decode(sa, raw_r) + per[stem]["stft_in"].append(multi_stft(ad.mean(0), ac.mean(0))) + per[stem]["stft_out"].append(multi_stft(ar.mean(0), ac.mean(0))) + + clean_mix = ac if clean_mix is None else clean_mix + ac + deg_mix = ad if deg_mix is None else deg_mix + ad + rest_mix = ar if rest_mix is None else rest_mix + ar + + if pi < cfg.demo_pairs: + d = demo_root / pid; d.mkdir(parents=True, exist_ok=True) + sf.write(d / f"{stem}_1_degraded.wav", ad.T.numpy(), sr) + sf.write(d / f"{stem}_2_clean.wav", ac.T.numpy(), sr) + sf.write(d / f"{stem}_3_restored.wav", ar.T.numpy(), sr) + + if clat is not None: + frmx["clean"].append(clat); frmx["in"].append(dlat); frmx["out"].append(rlat) + if clean_mix is not None: + remix["stft_in"].append(multi_stft(deg_mix.mean(0), clean_mix.mean(0))) + remix["stft_out"].append(multi_stft(rest_mix.mean(0), clean_mix.mean(0))) + if pi < cfg.demo_pairs: + d = demo_root / pid + sf.write(d / "MIX_1_degraded.wav", deg_mix.T.numpy(), sr) + sf.write(d / "MIX_2_clean.wav", clean_mix.T.numpy(), sr) + sf.write(d / "MIX_3_restored.wav", rest_mix.T.numpy(), sr) + + # report + print(f"\n=== eval epoch {epoch} (val pairs={len(pairs)}) ===") + print(f"{'stem':7} {'n':>4} | cos_in->out | relL2 | energy | stft_in->out") + out = {} + for s in cfg.stems: + if not per[s]["cos_out"]: + continue + m = {k: float(np.mean(v)) for k, v in per[s].items()} + out[s] = m + print(f"{s:7} {len(per[s]['cos_out']):>4} | {m['cos_in']:.3f}->{m['cos_out']:.3f} | " + f"{m['relL2_out']:.3f} | {m['energy_out']:.3f} | {m['stft_in']:.3f}->{m['stft_out']:.3f}") + if remix["stft_out"]: + ri, ro = float(np.mean(remix["stft_in"])), float(np.mean(remix["stft_out"])) + out["remix"] = {"stft_in": ri, "stft_out": ro} + print(f"{'REMIX':7} {len(remix['stft_out']):>4} | full-mix multi-STFT degraded={ri:.3f} -> restored={ro:.3f}") + # latent-FAD: distributional distance to clean (no decode) โ€” sees realism/dullness the + # reference-matching multi-STFT cannot. degraded(in) -> restored(out), lower = closer to clean. + try: + print(f"{'stem':7} | latent-FAD in->out (lower=closer)") + for s in cfg.stems: + if not flat["out"][s]: + continue + cl = fad.latent_set(flat["clean"][s]) + fi = fad.fad(cl, fad.latent_set(flat["in"][s])); fo = fad.fad(cl, fad.latent_set(flat["out"][s])) + out.setdefault(s, {}).update(fad_in=fi, fad_out=fo) + print(f"{s:7} | {fi:.3f}->{fo:.3f}") + if frmx["out"]: + clr = fad.latent_set(frmx["clean"]) + fi = fad.fad(clr, fad.latent_set(frmx["in"])); fo = fad.fad(clr, fad.latent_set(frmx["out"])) + out["fad_remix"] = {"in": fi, "out": fo} + print(f"{'FAD-REMIX':12} | {fi:.3f}->{fo:.3f} (distributional; complements REMIX-STFT)") + except Exception as e: + print(f"[eval] latent-FAD skipped: {e}") + print(f"(demo audio -> {demo_root})") + return out diff --git a/restoflow/fad.py b/restoflow/fad.py new file mode 100644 index 0000000000000000000000000000000000000000..ac7530ad49edf0b56f4c783b1189b85c9ea7e6c5 --- /dev/null +++ b/restoflow/fad.py @@ -0,0 +1,122 @@ +"""FAD (Frechet Audio Distance) for SAME-latent stem restoration/generation. + +Backends (same Frechet core): + - CHEAP : SAME latents directly (frame-level [256] vectors) โ€” NO decode, NO extra model. + "The latent is the mirror of the audio" (user): SAME's encoder is already a learned + audio representation, so Frechet on latents should track real FAD if discriminative. + - REAL : decode -> embedder. PREFERRED = laion **CLAP music** (512-d/clip @ 48 kHz, the + modern music-FAD embedder, used by fadtk/Diff-A-Riff). VGGish kept as a secondary. + +FAD(A,B) = ||mu_A-mu_B||^2 + Tr(Sig_A + Sig_B - 2 (Sig_A Sig_B)^{1/2}). Lower = test set's +distribution is closer to the real (clean) set's โ€” credits valid-but-different (one-to-many) +outputs that per-sample STFT/cosine wrongly penalize. Computed PER STEM and for the REMIX. +""" +from __future__ import annotations +import numpy as np +import torch +from scipy import linalg + + +# ---------------- Frechet core ---------------- +def gaussian(X: np.ndarray, shrink=None): + """X [N,D] -> (mu[D], cov[D,D]). Ledoit-Wolf shrinkage when samples are scarce vs dim + (auto when N < 4D) so the 512-d CLAP cov stays PSD/invertible on small per-stem sets.""" + X = np.asarray(X, dtype=np.float64) + mu = X.mean(0) + if shrink is None: + shrink = X.shape[0] < 4 * X.shape[1] + if shrink: + from sklearn.covariance import LedoitWolf + cov = LedoitWolf().fit(X).covariance_ + else: + cov = np.cov(X, rowvar=False) + return mu, np.atleast_2d(cov) + + +def frechet(mu1, cov1, mu2, cov2, eps=1e-6) -> float: + diff = mu1 - mu2 + covmean, _ = linalg.sqrtm(cov1 @ cov2, disp=False) + if not np.isfinite(covmean).all(): + off = np.eye(cov1.shape[0]) * eps + covmean = linalg.sqrtm((cov1 + off) @ (cov2 + off)) + if np.iscomplexobj(covmean): + covmean = covmean.real + return float(diff @ diff + np.trace(cov1) + np.trace(cov2) - 2 * np.trace(covmean)) + + +def fad(A: np.ndarray, B: np.ndarray, shrink=None) -> float: + """A=real embeddings [Na,D], B=test embeddings [Nb,D].""" + return frechet(*gaussian(A, shrink), *gaussian(B, shrink)) + + +# ---------------- CHEAP backend: SAME latents ---------------- +def latent_emb(latent: torch.Tensor) -> np.ndarray: + """[256,T] -> [T,256] frame-level embeddings (each frame = one sample).""" + return latent.float().t().cpu().numpy() + + +def latent_set(latents) -> np.ndarray: + return np.concatenate([latent_emb(l) for l in latents], 0) + + +# ---------------- REAL backend: CLAP music (preferred) ---------------- +class CLAP: + """laion CLAP music embedder โ€” 512-d per clip @ 48 kHz.""" + DEFAULT_CKPT = ("/home/ksoil/.local/lib/python3.10/site-packages/fadtk/" + ".model-checkpoints/music_audioset_epoch_15_esc_90.14.pt") + + def __init__(self, device="cpu", ckpt=None): + import laion_clap, torchaudio + self.m = laion_clap.CLAP_Module(enable_fusion=False, amodel="HTSAT-base", device=device) + self.m.load_ckpt(ckpt or self.DEFAULT_CKPT, verbose=False) + self.device, self._ta, self.rs = device, torchaudio, {} + + def _to48k_mono(self, audio, sr): + wav = audio.mean(0) if audio.dim() == 2 else audio # [S] + if sr != 48000: + if sr not in self.rs: + self.rs[sr] = self._ta.transforms.Resample(sr, 48000).to(wav.device) + wav = self.rs[sr](wav) + return wav + + @torch.inference_mode() + def embed_set(self, audios, sr, bs=128) -> np.ndarray: + """Faithful to fadtk/BABE-2: ONE 512-d CLAP-music embedding per input clip (no custom + windowing). Validity comes from supplying enough clips, not from re-segmenting here.""" + wavs = [self._to48k_mono(a, sr) for a in audios] + L = max(w.numel() for w in wavs) + wavs = [torch.nn.functional.pad(w, (0, L - w.numel())) for w in wavs] + out = [] + for i in range(0, len(wavs), bs): + x = torch.stack(wavs[i:i + bs]).to(self.device) + e = self.m.get_audio_embedding_from_data(x=x, use_tensor=True) + out.append(e.detach().cpu().numpy().astype(np.float64)) + return np.concatenate(out, 0) # [N_clips, 512] + + +# ---------------- REAL backend: VGGish (secondary) ---------------- +class VGGish: + """VGGish embedder (16 kHz mono -> 128-d per ~0.96 s).""" + def __init__(self, device="cpu"): + from torchvggish import vggish + import torchaudio + self.m = vggish(postprocess=False).to(device).eval() + self.device, self._ta, self.rs = device, torchaudio, {} + + @torch.inference_mode() + def embed(self, audio, sr): + from torchvggish import vggish_input + wav = audio.mean(0, keepdim=True) if audio.dim() == 2 else audio[None] + if sr != 16000: + if sr not in self.rs: + self.rs[sr] = self._ta.transforms.Resample(sr, 16000).to(wav.device) + wav = self.rs[sr](wav) + ex = vggish_input.waveform_to_examples(wav.squeeze(0).cpu().numpy(), 16000, return_tensor=True) + if ex.numel() == 0: + return np.zeros((0, 128)) + return self.m(ex.to(self.device)).cpu().numpy().astype(np.float64) + + def embed_set(self, audios, sr): + cs = [self.embed(a, sr) for a in audios] + cs = [c for c in cs if len(c)] + return np.concatenate(cs, 0) if cs else np.zeros((0, 128)) diff --git a/restoflow/fad_eval.py b/restoflow/fad_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..ee8a6de7eb064a8a20b609f61a67ba2fd4035bcd --- /dev/null +++ b/restoflow/fad_eval.py @@ -0,0 +1,135 @@ +"""Validate CHEAP latent-FAD vs REAL CLAP-music-FAD, per stem + REMIX. + +For a sample of val pairs, builds 4 condition sets vs clean: degraded, det-restored +(restorer_attn_v1), flowattn-restored (scout_rest_flowattn), plus a clean-split FLOOR. +Reports FAD(clean, X) for both backends. If the two backends RANK the conditions the same, +the cheap latent-FAD is a valid stand-in (no decode, no embedder). + +Run: python -m restoflow.fad_eval --n-pairs 150 --device cuda +""" +from __future__ import annotations +import argparse, random +from collections import defaultdict +import numpy as np, torch + +from .config import Cfg, STEMS, STEM_ID +from . import data as D, flow as Fl, eval as E, fad +from .model import AttnRestorer, MixAttnCondFlow + +BASE = "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra" +DET = f"{BASE}/restoflow_runs/restorer_attn_v1" +FLOW = f"{BASE}/restoflow_runs/scout_rest_flowattn" +CONDS = ["degraded", "det", "flowattn"] + + +def _load_det(dev): + ck = torch.load(f"{DET}/ckpt_best.pt", map_location="cpu"); c = ck["cfg"] + m = AttnRestorer(256, c["hidden"], c["depth"], n_stems=len(STEMS), stem_emb=c["stem_emb_dim"], + use_mix=c["use_mix"], heads=c.get("heads", 4)).to(dev).eval() + m.load_state_dict(ck["model"]); return m, torch.load(f"{DET}/norm_stats.pt", map_location="cpu") + + +def _load_flow(dev): + ck = torch.load(f"{FLOW}/ckpt_best.pt", map_location="cpu"); c = ck["cfg"] + m = MixAttnCondFlow(256, c["hidden"], c["depth"], n_stems=len(STEMS), stem_emb=c["stem_emb_dim"], + heads=c.get("heads", 8), use_mix=c["use_mix"]).to(dev).eval() + m.load_state_dict(ck["model"]); return m, torch.load(f"{FLOW}/norm_stats.pt", map_location="cpu"), c.get("sigma", 0.3) + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--n-pairs", type=int, default=150) + p.add_argument("--device", default="cuda") + p.add_argument("--steps", type=int, default=40) + p.add_argument("--no-clap", action="store_true") + a = p.parse_args() + dev = a.device + random.seed(0); torch.manual_seed(0) + + cfg = Cfg(val_cache_roots=(f"{BASE}/demucs_results_val_full",), use_mix=True, model_kind="attn", + T=32, device=dev, sample_steps=a.steps) + sa, sr = E.load_same(cfg) + det, dstats = _load_det(dev) + fl, fstats, sigma = _load_flow(dev) + print(f"[fad] models loaded; flow sigma={sigma}; sr={sr}") + + items = D._index(cfg, cfg.val_cache_roots, "val") + by_pair = defaultdict(list) + for it in items: + by_pair[it["pair_id"]].append(it) + pairs = list(by_pair.items()); random.shuffle(pairs) + freq = defaultdict(int) # front-load rare-stem (bass/drums) pairs + for _, its in pairs: + for it in its: freq[it["stem"]] += 1 + pairs.sort(key=lambda kv: min(freq[it["stem"]] for it in kv[1])) + pairs = pairs[:a.n_pairs] + print(f"[fad] {len(pairs)} pairs, {sum(len(v) for _,v in pairs)} stem-instances") + + # collectors: lat[cond][stem] = list[256,T]; aud[cond][stem] = list[2,S]; remix per pair + lat = {c: defaultdict(list) for c in ["clean"] + CONDS} + aud = {c: defaultdict(list) for c in ["clean"] + CONDS} + rlat = {c: [] for c in ["clean"] + CONDS} + raud = {c: [] for c in ["clean"] + CONDS} + + def destd(std, stats, stem): return E._destd(std[0], stats, stem) + + with torch.inference_mode(): + for pi, (pid, its) in enumerate(pairs): + psum = {c: None for c in ["clean"] + CONDS}; pasum = {c: None for c in ["clean"] + CONDS} + for it in its: + stem = it["stem"]; sid = torch.tensor([STEM_ID[stem]], device=dev) + raw_c = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) + raw_d = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) + raw_m = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) + # det normalization + dmu, dsd = dstats[stem]["mu"], dstats[stem]["sd"]; mm = dstats["__mix__"] + std_d = ((raw_d - dmu[:, None]) / dsd[:, None]).to(dev)[None] + std_mix = ((raw_m - mm["mu"][:, None]) / mm["sd"][:, None]).to(dev)[None] + raw_det = destd(det(std_d, sid, std_mix), dstats, stem) + # flow normalization (own stats) + fmu, fsd = fstats[stem]["mu"], fstats[stem]["sd"]; fm = fstats["__mix__"] + fstd_d = ((raw_d - fmu[:, None]) / fsd[:, None]).to(dev)[None] + fstd_mix = ((raw_m - fm["mu"][:, None]) / fm["sd"][:, None]).to(dev)[None] + fstd_r = Fl.sample_mix(fl, fstd_d, fstd_mix, sid, cfg.sample_steps, sigma) + raw_flow = destd(fstd_r, fstats, stem) + raws = {"clean": raw_c.to(dev), "degraded": raw_d.to(dev), "det": raw_det, "flowattn": raw_flow} + for c in ["clean"] + CONDS: + lat[c][stem].append(raws[c].cpu()) + au = E._decode(sa, raws[c]) + aud[c][stem].append(au) + psum[c] = raws[c] if psum[c] is None else psum[c] + raws[c] + pasum[c] = au if pasum[c] is None else pasum[c] + au + for c in ["clean"] + CONDS: + if psum[c] is not None: + rlat[c].append(psum[c].cpu()); raud[c].append(pasum[c]) + if (pi + 1) % 25 == 0: print(f" ...{pi+1}/{len(pairs)} pairs") + + clap = None if a.no_clap else fad.CLAP(device=dev) + + def report(title, lat_sets, aud_sets): + # lat_sets/aud_sets: dict cond -> list (latents [256,T] / audios [2,S]) + n = len(lat_sets["clean"]); half = n // 2 + clean_lat = fad.latent_set(lat_sets["clean"]) + floor_lat = (fad.fad(fad.latent_set(lat_sets["clean"][:half]), fad.latent_set(lat_sets["clean"][half:])) + if half >= 1 else float("nan")) + row = {c: fad.fad(clean_lat, fad.latent_set(lat_sets[c])) for c in CONDS} + print(f"\n[{title}] LATENT-FAD (cheap) floor(clean-split)={floor_lat:.3f}") + print(" " + " ".join(f"{c}={row[c]:.3f}" for c in CONDS)) + if clap is not None: + clean_au = clap.embed_set(aud_sets["clean"], sr) + floor_au = (fad.fad(clap.embed_set(aud_sets["clean"][:half], sr), clap.embed_set(aud_sets["clean"][half:], sr)) + if half >= 1 else float("nan")) + rowc = {c: fad.fad(clean_au, clap.embed_set(aud_sets[c], sr)) for c in CONDS} + print(f"[{title}] CLAP-FAD (real) floor(clean-split)={floor_au:.3f}") + print(" " + " ".join(f"{c}={rowc[c]:.3f}" for c in CONDS)) + + for stem in STEMS: + if lat["clean"][stem]: + report(f"stem={stem} n={len(lat['clean'][stem])}", + {c: lat[c][stem] for c in ["clean"] + CONDS}, + {c: aud[c][stem] for c in ["clean"] + CONDS}) + report(f"REMIX n={len(rlat['clean'])}", rlat, raud) + + +if __name__ == "__main__": + main() diff --git a/restoflow/fill_clean.py b/restoflow/fill_clean.py new file mode 100644 index 0000000000000000000000000000000000000000..ca3e65b41a972c6d2c11994e8db02d91e2e45979 --- /dev/null +++ b/restoflow/fill_clean.py @@ -0,0 +1,238 @@ +"""Incremental clean-stem latent fill. + +The original cache only wrote clean/degraded latents for `class == "restoration"` +(demucs.py:1242). Generation-class stems (clean present, degraded gutted/absent) +got NO clean latent -> the generator was starved of ~66% of clean targets. + +This fills ONLY the missing clean latents, ONLY for affected pairs, by re-separating +the CLEAN source mix with the same htdemucs path and SAME-encoding the needed stems. +It does not touch degraded (generation cases don't need a degraded target latent). + +Reuses demucs.py's tested separation + encode primitives so latents match exactly +(same normalization, segment_size, dtype). Skip-existing -> resumable / incremental. + +Run per GPU shard, e.g.: + EE=/home/ksoil/.conda/envs/ksoil_encoders/bin/python + $EE -m restoflow.fill_clean --output-root demucs_results_all \ + --clean-dir $SRC/train/clean --gpu 0 --num-shards 3 --shard 0 \ + --results-out fill_clean_train_s0.csv +Then merge into metadata.csv: + $EE -m restoflow.fill_clean --merge --output-root demucs_results_all \ + --results-glob 'fill_clean_train_s*.csv' +""" +from __future__ import annotations + +import argparse +import csv +import glob as globmod +import os +import sys +from collections import defaultdict +from pathlib import Path + +BASE = Path(__file__).resolve().parent.parent + +# Import the local demucs.py UNDER A DIFFERENT NAME. Importing it as `demucs` +# would shadow the real `demucs` package that the htdemucs checkpoint needs at +# unpickle time (-> "No module named 'demucs.htdemucs'"). +import importlib.util as _ilu # noqa: E402 +_spec = _ilu.spec_from_file_location("demucs_local", BASE / "demucs.py") +D = _ilu.module_from_spec(_spec) +_spec.loader.exec_module(D) + +RESULT_COLS = ["pair_id", "stem", "clean_latent_path", "latent_dim", "n_frames", + "frame_rate", "duration_s"] + + +def affected_pairs(output_root: Path) -> dict[str, list[str]]: + """pair_id -> sorted list of stems with clean present but no clean latent.""" + need: dict[str, list[str]] = defaultdict(list) + with open(output_root / "metadata.csv") as f: + for r in csv.DictReader(f): + if r.get("clean_present") == "1" and not r.get("clean_latent_path"): + need[r["pair_id"]].append(r["stem"]) + return {pid: sorted(set(stems)) for pid, stems in need.items()} + + +def build_args() -> argparse.Namespace: + """A demucs Namespace with all defaults, suited for single-pass clean separate+encode.""" + args = D.parse_args([]) + args.same_encode = True + args.persistent_demucs = True + args.quiet_demucs = True + args.demucs_segment_size = "3" + return args + + +def run_shard(output_root: Path, clean_dir: Path, gpu: int, num_shards: int, + shard: int, results_out: Path) -> None: + os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu) + D.configure_process_env() + + args = build_args() + args.output_root = output_root + args.clean_dir = clean_dir + + need = affected_pairs(output_root) + pairs = sorted(need) + # deterministic shard by index + pairs = [p for i, p in enumerate(pairs) if i % num_shards == shard] + print(f"[shard {shard}/{num_shards} gpu{gpu}] {len(pairs)} affected pairs", flush=True) + + from audio_separator.separator.architectures import demucs_separator + with D.suppress_output(True): + separator = D.make_separator(args) + separator.load_model(model_filename=args.model) + model_instance = D.init_persistent_demucs(separator) + if model_instance is None: + raise RuntimeError("persistent htdemucs model unavailable") + with D.suppress_output(True): + same_ctx = D.load_same_model(args) + print(f"[shard {shard}] models ready (SAME {'fp16' if same_ctx['half'] else 'fp32'})", flush=True) + + # resume: skip already-written rows + done: set[tuple[str, str]] = set() + if results_out.exists(): + with open(results_out) as f: + for r in csv.DictReader(f): + done.add((r["pair_id"], r["stem"])) + + new = not results_out.exists() + fout = open(results_out, "a", newline="") + w = csv.DictWriter(fout, fieldnames=RESULT_COLS) + if new: + w.writeheader() + + import torch + n_enc = 0 + for i, pid in enumerate(D.tqdm_iter(pairs, args, total=len(pairs), desc=f"fill:{shard}", + unit="pair", dynamic_ncols=True, leave=True)): + stems_need = need[pid] + # skip stems already present on disk or already recorded + todo = [] + for s in stems_need: + lp = D.latent_output_path(args, pid, "clean", s) + if (pid, s) in done: + continue + if lp.exists(): + # already encoded (prior/partial run) but not yet recorded -> emit a row + z = torch.load(str(lp), map_location="cpu") + w.writerow({ + "pair_id": pid, "stem": s, + "clean_latent_path": str(lp.relative_to(output_root)), + "latent_dim": int(z.shape[0]), "n_frames": int(z.shape[1]), + "frame_rate": round(same_ctx["sr"] * int(z.shape[1]) / (same_ctx["sr"] * 3), 4), + "duration_s": 3.0, + }) + done.add((pid, s)) + continue + todo.append(s) + if not todo: + continue + + wav = clean_dir / f"{pid}.wav" + if not wav.exists(): + print(f"[shard {shard}] WARN missing clean source {wav}", flush=True) + continue + + try: + with torch.inference_mode(), D.suppress_output(True): + mix = model_instance.prepare_mix(str(wav)) + source = model_instance.demix_demucs(mix) + except Exception as e: # noqa: BLE001 + print(f"[shard {shard}] ERROR separate {pid}: {e}", flush=True) + continue + + n = len(source) + smap = (demucs_separator.DEMUCS_2_SOURCE_MAPPER if n == 2 else + demucs_separator.DEMUCS_6_SOURCE_MAPPER if n == 6 else + demucs_separator.DEMUCS_4_SOURCE_MAPPER) + idx_of = {D.normalize_token(k): v for k, v in smap.items()} + + for s in todo: + if s not in idx_of: + continue + stem_TC = D.normalized_stem_array(model_instance, source[idx_of[s]].T) + latent, _, _ = D.same_encode_stem(same_ctx, stem_TC, args.sample_rate) + lp = D.latent_output_path(args, pid, "clean", s) + D.save_latent_pt(lp, latent, args.same_latent_dtype) + n_frames = int(latent.shape[1]) + n_samples = int(stem_TC.shape[0]) + fr = round(same_ctx["sr"] * n_frames / n_samples, 4) if n_samples else 0.0 + w.writerow({ + "pair_id": pid, "stem": s, + "clean_latent_path": str(lp.relative_to(output_root)), + "latent_dim": int(latent.shape[0]), "n_frames": n_frames, + "frame_rate": fr, + "duration_s": round(n_samples / same_ctx["sr"], 4) if same_ctx["sr"] else 0.0, + }) + n_enc += 1 + if (i + 1) % 200 == 0: + fout.flush() + model_instance.clear_gpu_cache() + fout.close() + print(f"[shard {shard}] DONE encoded {n_enc} clean latents", flush=True) + + +def merge(output_root: Path, results_glob: str) -> None: + """Fold filled clean_latent_path (+dims) back into metadata.csv.""" + upd: dict[tuple[str, str], dict] = {} + for fp in sorted(globmod.glob(str(output_root / results_glob)) or globmod.glob(results_glob)): + with open(fp) as f: + for r in csv.DictReader(f): + upd[(r["pair_id"], r["stem"])] = r + print(f"merge: {len(upd)} filled (pair,stem) rows") + + meta = output_root / "metadata.csv" + rows = [] + with open(meta) as f: + rd = csv.DictReader(f) + fields = rd.fieldnames + for r in rd: + key = (r["pair_id"], r["stem"]) + if key in upd and not r.get("clean_latent_path"): + u = upd[key] + r["clean_latent_path"] = u["clean_latent_path"] + if not r.get("latent_dim"): + r["latent_dim"] = u["latent_dim"] + if not r.get("n_frames"): + r["n_frames"] = u["n_frames"] + if not r.get("frame_rate"): + r["frame_rate"] = u["frame_rate"] + rows.append(r) + + bak = meta.with_suffix(".csv.prefill.bak") + if not bak.exists(): + os.replace(meta, bak) + src = bak + else: + src = meta # already backed up; just rewrite + # (we already consumed meta into rows; write fresh) + with open(meta, "w", newline="") as f: + wr = csv.DictWriter(f, fieldnames=fields) + wr.writeheader() + wr.writerows(rows) + filled = sum(1 for r in rows if r["clean_latent_path"]) + print(f"merge: wrote {meta} (backup {src.name}); rows with clean_latent_path now: {filled}") + + +def main() -> None: + p = argparse.ArgumentParser() + p.add_argument("--output-root", type=Path, required=True) + p.add_argument("--clean-dir", type=Path) + p.add_argument("--gpu", type=int, default=0) + p.add_argument("--num-shards", type=int, default=1) + p.add_argument("--shard", type=int, default=0) + p.add_argument("--results-out", type=Path) + p.add_argument("--merge", action="store_true") + p.add_argument("--results-glob", default="fill_clean_*.csv") + a = p.parse_args() + if a.merge: + merge(a.output_root, a.results_glob) + else: + assert a.clean_dir and a.results_out, "need --clean-dir and --results-out" + run_shard(a.output_root, a.clean_dir, a.gpu, a.num_shards, a.shard, a.results_out) + + +if __name__ == "__main__": + main() diff --git a/restoflow/fill_restored.py b/restoflow/fill_restored.py new file mode 100644 index 0000000000000000000000000000000000000000..778531a8380c114b0e6a5d2787d0a8bffef36c20 --- /dev/null +++ b/restoflow/fill_restored.py @@ -0,0 +1,71 @@ +"""Incrementally FILL latents/restored/{stem}.pt by running the v4b restorer over cached +degraded stem latents (+ degraded_mix). Latent->latent, no decode -> fast. Skips existing. + +These restored latents are the realistic CONTEXT for the generator: at inference we never have +clean stems, only v4b restorations. Run: python -m restoflow.fill_restored --root demucs_results_all +""" +from __future__ import annotations +import argparse, csv, os +from pathlib import Path +import torch + +from .config import Cfg, STEM_ID, STEMS +from .model import DetRestorer, AttnRestorer + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--root", required=True) + ap.add_argument("--ckpt", default="restoflow_runs/restorer_attn_distvar_w003/ckpt_best.pt") + ap.add_argument("--stats", default="restoflow_runs/restorer_attn_distvar_w003/norm_stats.pt") + ap.add_argument("--device", default="cuda") + ap.add_argument("--overwrite", action="store_true", help="regenerate even if restored/{stem}.pt exists") + a = ap.parse_args() + c = Cfg(); dev = a.device + BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") + root = BASE / a.root + + ck = torch.load(BASE / a.ckpt, map_location="cpu"); cfg = ck["cfg"] + saved = {k: v for k, v in cfg.items() if k in ("latent_dim", "hidden", "depth", "stem_emb_dim", "use_mix")} + if cfg.get("model_kind") == "attn": # advramp/cotrain/w003 are AttnRestorer; v4b was DetRestorer + model = AttnRestorer(saved.get("latent_dim", 256), saved["hidden"], saved["depth"], n_stems=len(STEMS), + stem_emb=saved["stem_emb_dim"], use_mix=saved["use_mix"], heads=cfg.get("heads", 4)) + else: + model = DetRestorer(saved.get("latent_dim", 256), saved["hidden"], saved["depth"], + n_stems=len(STEMS), stem_emb=saved["stem_emb_dim"], use_mix=saved["use_mix"]) + model.load_state_dict(ck["model"]); model.eval().to(dev) + stats = torch.load(BASE / a.stats, map_location="cpu") + mmu, msd = stats["__mix__"]["mu"], stats["__mix__"]["sd"] + print(f"[fill] v4b loaded; root={root}") + + rows = list(csv.DictReader(open(root / "metadata.csv"))) + made = skipped = nomix = 0 + with torch.inference_mode(): + for r in rows: + stem = r.get("stem") + if stem not in STEMS or r.get("deg_present") != "1" or not r.get("deg_latent_path"): + continue + pair = r["pair_id"] + out = root / pair / "latents" / "restored" / f"{stem}.pt" + if out.exists() and not a.overwrite: + skipped += 1; continue + mix_p = root / pair / "latents" / "degraded_mix.pt" + if not mix_p.exists(): + nomix += 1; continue + deg = torch.load(root / r["deg_latent_path"], map_location="cpu").float() + mix = torch.load(mix_p, map_location="cpu").float() + mu, sd = stats[stem]["mu"], stats[stem]["sd"] + std_d = ((deg - mu[:, None]) / sd[:, None]).to(dev)[None] + std_m = ((mix - mmu[:, None]) / msd[:, None]).to(dev)[None] + sid = torch.tensor([STEM_ID[stem]], device=dev) + std_r = model(std_d, sid, std_m)[0].cpu() + raw_r = std_r * sd[:, None] + mu[:, None] + out.parent.mkdir(parents=True, exist_ok=True) + torch.save(raw_r.half(), out); made += 1 + if made % 2000 == 0: + print(f" made={made} skipped={skipped}") + print(f"[fill] done. made={made} skipped(existing)={skipped} no_mix={nomix}") + + +if __name__ == "__main__": + main() diff --git a/restoflow/final_eval.py b/restoflow/final_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..308ae57224342c49e618afaf8ae5d5d5ce156592 --- /dev/null +++ b/restoflow/final_eval.py @@ -0,0 +1,52 @@ +"""Correct final eval: load a trained checkpoint and score it over the FULL val cache +with stratified (all-stem) coverage. Uses the run's TRAIN norm stats (not rebuilt). + + python -m restoflow.final_eval --ckpt restoflow_runs/v3_big/ckpt_best.pt \ + --val-cache-roots demucs_results_val_full --eval-max-pairs 0 +""" +from __future__ import annotations +import argparse +from dataclasses import replace +from pathlib import Path +import torch + +from .config import Cfg +from . import eval as E +from .model import CondFlow, DetRestorer + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--ckpt", required=True) + p.add_argument("--val-cache-roots", default="demucs_results_val_full") + p.add_argument("--eval-max-pairs", type=int, default=0, help="0 = all val pairs") + p.add_argument("--device", default="cuda") + p.add_argument("--demo-pairs", type=int, default=8) + a = p.parse_args() + + ck = torch.load(a.ckpt, map_location="cpu") + base = Cfg() + saved = {k: v for k, v in ck["cfg"].items() if hasattr(base, k)} + cfg = replace(base, **saved) + run_dir = Path(a.ckpt).parent + cfg = replace(cfg, + val_cache_roots=tuple(x for x in a.val_cache_roots.split(",") if x), + eval_max_pairs=a.eval_max_pairs, device=a.device, demo_pairs=a.demo_pairs, + out_dir=str(run_dir), stats_path=str(run_dir / "norm_stats.pt")) + + stats = torch.load(cfg.stats_file(), map_location="cpu") + if cfg.model_kind == "det": + model = DetRestorer(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim, use_mix=cfg.use_mix) + else: + model = CondFlow(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim) + model.load_state_dict(ck["model"]); model = model.to(cfg.device) + print(f"[final_eval] ckpt={a.ckpt} epoch={ck.get('epoch')} {cfg.model_kind} " + f"params={model.num_params()/1e6:.2f}M val={cfg.val_cache_roots} budget={a.eval_max_pairs or 'ALL'}") + sa, sr = E.load_same(cfg) + E.evaluate(cfg, model, stats, sa, sr, epoch=9999) + + +if __name__ == "__main__": + main() diff --git a/restoflow/flow.py b/restoflow/flow.py new file mode 100644 index 0000000000000000000000000000000000000000..a2b494deecc7e6ac8d750862a7475d378b8ed958 --- /dev/null +++ b/restoflow/flow.py @@ -0,0 +1,68 @@ +"""Flow-matching objective + ODE sampler for the deg-anchored restoration bridge. + +Noisy-start rectified flow: x0 = degraded + sigma*eps, x1 = clean, x_t=(1-t)x0+t x1, +target velocity v = x1 - x0. Diversity comes from eps; sampling integrates v from a +degraded-anchored start -> a plausible clean (the restoration trajectory). +""" +from __future__ import annotations +import torch +import torch.nn.functional as F + + +def fm_loss(model, deg, clean, stem_id, sigma: float): + B = deg.shape[0] + eps = torch.randn_like(clean) + x0 = deg + sigma * eps + t = torch.rand(B, device=deg.device) + tt = t[:, None, None] + z_t = (1 - tt) * x0 + tt * clean + v_target = clean - x0 + return F.mse_loss(model(z_t, t, deg, stem_id), v_target) + + +@torch.inference_mode() +def sample(model, deg, stem_id, steps: int, sigma: float, generator=None): + eps = torch.randn(deg.shape, generator=generator, device=deg.device) + z = deg + sigma * eps + dt = 1.0 / steps + for i in range(steps): + t = torch.full((deg.shape[0],), i * dt, device=deg.device) + z = z + dt * model(z, t, deg, stem_id) + return z + + +def fm_loss_mix(model, deg, clean, mix, stem_id, sigma: float): + """Deg-anchored flow loss WITH mix conditioning (model takes z_t,t,deg,stem_id,mix).""" + B = deg.shape[0] + eps = torch.randn_like(clean) + x0 = deg + sigma * eps + t = torch.rand(B, device=deg.device) + tt = t[:, None, None] + z_t = (1 - tt) * x0 + tt * clean + v_target = clean - x0 + return F.mse_loss(model(z_t, t, deg, stem_id, mix), v_target) + + +@torch.inference_mode() +def sample_mix(model, deg, mix, stem_id, steps: int, sigma: float, generator=None): + eps = torch.randn(deg.shape, generator=generator, device=deg.device) + z = deg + sigma * eps + dt = 1.0 / steps + for i in range(steps): + t = torch.full((deg.shape[0],), i * dt, device=deg.device) + z = z + dt * model(z, t, deg, stem_id, mix) + return z + + +# ---- cheap latent metrics (validated proxies for decoded-audio distance) ---- +def latent_cos_dist(pred, clean): # best proxy (within-clip Spearman ~0.94) + p, c = pred.flatten(1), clean.flatten(1) + return (1 - F.cosine_similarity(p, c, dim=1)).mean() + +def latent_relL2(pred, clean): + p, c = pred.flatten(1), clean.flatten(1) + return ((p - c).norm(dim=1) / (c.norm(dim=1) + 1e-9)).mean() + +def energy_ratio(pred, clean): + p, c = pred.flatten(1), clean.flatten(1) + return (p.norm(dim=1) / (c.norm(dim=1) + 1e-9)).mean() diff --git a/restoflow/gen.py b/restoflow/gen.py new file mode 100644 index 0000000000000000000000000000000000000000..531a0dcbaf80b3ce962a8aca8b01ba2b0a22b36c --- /dev/null +++ b/restoflow/gen.py @@ -0,0 +1,552 @@ +"""Conditional stem GENERATION (accompaniment) in SAME-L latent space. + +Probe-validated (GENERATION_PLAN.md sec 11): context = sum of OTHER clean stem latents +strongly determines bass/drums (cos 0.33/0.42); it is one-to-many (relL2 ~0.72) -> GENERATIVE. + +Model: conditional flow-matching from NOISE -> target stem latent, conditioned on the +Sigma-context + stem-id (reuses model.CondFlow; its `deg` slot = context). Classifier-free +guidance via context dropout. NO reference-matching loss (one-to-many). Trains on CLEAN +leave-one-out (no degraded data needed). + +Run: python -m restoflow.gen --epochs 200 --target-stems bass,drums +""" +from __future__ import annotations +import argparse, glob, json, os, random, time +from collections import defaultdict +from pathlib import Path +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch.utils.data import Dataset, DataLoader + +from .config import Cfg, STEMS, STEM_ID +from .model import CondFlow, AttnCondFlow, XAttnCondFlow +from . import eval as E # load_same, multi_stft, _decode +from . import distloss as DL +from . import fad + +CTX_K = len(STEMS) - 1 # max context (other) stems per item, for the cross-attention generator + + +# ---------------- data: clean leave-one-out ---------------- +def _fitT(x, T): + t = x.shape[1] + if t == T: return x + return x[:, :T] if t > T else torch.cat([x, x[:, -1:].repeat(1, T - t)], 1) + + +def _src_path(root, pair, source, stem): + return f"{root}/{pair}/latents/{source}/{stem}.pt" + + +def index_clips(roots, target_stems, sources=("restored",)): + """Items for clean-target leave-one-out generation. Target = CLEAN held-out stem. + Context = sum of OTHER stems drawn from `sources` (realistic: 'restored'/'degraded', + NEVER clean at inference). Per item we store, for each source, the other-stem paths + that exist; a clip is kept only if >=1 source yields >=1 other. + """ + items = [] + for root in roots: + for cl in sorted(glob.glob(f"{root}/*/latents/clean")): + pair = Path(cl).parent.parent.name + clean = {s: f"{cl}/{s}.pt" for s in STEMS if os.path.exists(f"{cl}/{s}.pt")} + for tgt in target_stems: + if tgt not in clean: + continue + ctx_by_src = {} + for src in sources: + op = {s: _src_path(root, pair, src, s) for s in STEMS + if s != tgt and os.path.exists(_src_path(root, pair, src, s))} + if op: + ctx_by_src[src] = op + if not ctx_by_src: + continue + items.append({"pair": pair, "root": root, "target": tgt, + "tgt_path": clean[tgt], "ctx_by_src": ctx_by_src}) + return items + + +def _pick_ctx(it, rng): + """Pick one available context source (uniform over sources present for this item).""" + srcs = list(it["ctx_by_src"].keys()) + return srcs[rng.randrange(len(srcs))] if len(srcs) > 1 else srcs[0] + + +class GenData(Dataset): + def __init__(self, items, stats, T): + self.items, self.stats, self.T = items, stats, T + self.rng = random.Random(0) + + def __len__(self): return len(self.items) + + def __getitem__(self, i): + it = self.items[i]; s = it["target"] + tgt = _fitT(torch.load(it["tgt_path"], map_location="cpu").float(), self.T) + src = _pick_ctx(it, self.rng) + ctx = sum(_fitT(torch.load(p, map_location="cpu").float(), self.T) + for p in it["ctx_by_src"][src].values()) + tm, ts = self.stats["tgt"][s]; cm, cs = self.stats["ctx"][s] + return {"ctx": (ctx - cm[:, None]) / cs[:, None], + "tgt": (tgt - tm[:, None]) / ts[:, None], + "stem_id": torch.tensor(STEM_ID[s])} + + +def build_gen_stats(items, T, D, max_n=20000): + """Per-target-stem per-dim mean/std for clean target and for the context-sum + (context drawn the same way GenData draws it, so stats match the training distribution).""" + rng = random.Random(0) + sub = items if len(items) <= max_n else random.sample(items, max_n) + acc = {s: {"t_s": torch.zeros(D, dtype=torch.float64), "t_q": torch.zeros(D, dtype=torch.float64), + "c_s": torch.zeros(D, dtype=torch.float64), "c_q": torch.zeros(D, dtype=torch.float64), + "n": 0} for s in STEMS} + for it in sub: + s = it["target"] + tgt = _fitT(torch.load(it["tgt_path"], map_location="cpu").double(), T) + src = _pick_ctx(it, rng) + ctx = sum(_fitT(torch.load(p, map_location="cpu").double(), T) for p in it["ctx_by_src"][src].values()) + a = acc[s]; a["n"] += tgt.shape[1] + a["t_s"] += tgt.sum(1); a["t_q"] += (tgt * tgt).sum(1) + a["c_s"] += ctx.sum(1); a["c_q"] += (ctx * ctx).sum(1) + stats = {"tgt": {}, "ctx": {}} + for s, a in acc.items(): + if a["n"] == 0: + stats["tgt"][s] = (torch.zeros(D), torch.ones(D)); stats["ctx"][s] = (torch.zeros(D), torch.ones(D)); continue + tm = a["t_s"] / a["n"]; tv = (a["t_q"] / a["n"] - tm * tm).clamp_min(1e-8) + cm = a["c_s"] / a["n"]; cv = (a["c_q"] / a["n"] - cm * cm).clamp_min(1e-8) + stats["tgt"][s] = (tm.float(), tv.sqrt().float()) + stats["ctx"][s] = (cm.float(), cv.sqrt().float()) + return stats + + +# ---------------- per-stem context (cross-attention generator) ---------------- +class GenDataX(Dataset): + """Like GenData but keeps context stems SEPARATE (no sum): returns padded per-stem + context latents + their instrument ids + a present-mask, for XAttnCondFlow.""" + def __init__(self, items, stats, T, D): + self.items, self.stats, self.T, self.D = items, stats, T, D + self.rng = random.Random(0) + + def __len__(self): return len(self.items) + + def __getitem__(self, i): + it = self.items[i]; s = it["target"] + tm, ts = self.stats["tgt"][s] + tgt = _fitT(torch.load(it["tgt_path"], map_location="cpu").float(), self.T) + tgt = (tgt - tm[:, None]) / ts[:, None] + others = it["ctx_by_src"][_pick_ctx(it, self.rng)] + ctx = torch.zeros(CTX_K, self.D, self.T) + ids = torch.zeros(CTX_K, dtype=torch.long) + mask = torch.zeros(CTX_K, dtype=torch.bool) + for j, o in enumerate(sorted(others)): + x = _fitT(torch.load(others[o], map_location="cpu").float(), self.T) + sm, sd = self.stats["src"][o] + ctx[j] = (x - sm[:, None]) / sd[:, None] + ids[j] = STEM_ID[o]; mask[j] = True + return {"ctx_stems": ctx, "ctx_ids": ids, "ctx_mask": mask, + "tgt": tgt, "stem_id": torch.tensor(STEM_ID[s])} + + +def build_gen_stats_x(items, T, D, max_n=20000): + """Per-stem z-score stats: target stems normalized by their CLEAN distribution; context + stems by their AS-DRAWN (restored/degraded) per-instrument distribution.""" + rng = random.Random(0) + sub = items if len(items) <= max_n else random.sample(items, max_n) + z = lambda: {"s": torch.zeros(D, dtype=torch.float64), "q": torch.zeros(D, dtype=torch.float64), "n": 0} + tgt_acc = {s: z() for s in STEMS}; src_acc = {o: z() for o in STEMS} + for it in sub: + s = it["target"] + tgt = _fitT(torch.load(it["tgt_path"], map_location="cpu").double(), T) + a = tgt_acc[s]; a["n"] += tgt.shape[1]; a["s"] += tgt.sum(1); a["q"] += (tgt * tgt).sum(1) + for o, p in it["ctx_by_src"][_pick_ctx(it, rng)].items(): + x = _fitT(torch.load(p, map_location="cpu").double(), T) + b = src_acc[o]; b["n"] += x.shape[1]; b["s"] += x.sum(1); b["q"] += (x * x).sum(1) + + def fin(acc): + out = {} + for k, a in acc.items(): + if a["n"] == 0: + out[k] = (torch.zeros(D), torch.ones(D)); continue + m = a["s"] / a["n"]; v = (a["q"] / a["n"] - m * m).clamp_min(1e-8) + out[k] = (m.float(), v.sqrt().float()) + return out + return {"tgt": fin(tgt_acc), "src": fin(src_acc)} + + +def fm_loss_gen_x(model, ctx_stems, ctx_ids, ctx_mask, tgt, sid, cfg_drop, dist_kind="none", dist_w=0.0): + """Conditional flow matching with per-stem context. CFG: drop ALL context for a sample + (mask every real stem) so only the learned null token remains = unconditional.""" + B = tgt.shape[0] + if cfg_drop > 0: + keep = (torch.rand(B, device=tgt.device) >= cfg_drop)[:, None] + ctx_mask = ctx_mask & keep + x0 = torch.randn_like(tgt) + t = torch.rand(B, device=tgt.device); tt = t[:, None, None] + z = (1 - tt) * x0 + tt * tgt + v = tgt - x0 + vhat = model(z, t, ctx_stems, ctx_ids, ctx_mask, sid) + return F.mse_loss(vhat, v) + _endpoint_dist(z, tt, vhat, tgt, sid, dist_kind, dist_w) + + +@torch.inference_mode() +def sample_gen_x(model, ctx_stems, ctx_ids, ctx_mask, sid, steps, w, rescale=0.0, generator=None): + """Integrate noise->target with CFG. `rescale`>0 rescales the guided velocity back to the + conditional-pass norm each step (norm-preserving CFG) โ€” the principled fix for unbounded- + latent over-extrapolation that makes generated stems decode too loud (Sony arXiv:2402.01412).""" + B, _, D, T = ctx_stems.shape + z = torch.randn((B, D, T), generator=generator, device=ctx_stems.device) + null = torch.zeros_like(ctx_mask) + dt = 1.0 / steps + for i in range(steps): + t = torch.full((B,), i * dt, device=z.device) + vc = model(z, t, ctx_stems, ctx_ids, ctx_mask, sid) + if w != 1.0: + vu = model(z, t, ctx_stems, ctx_ids, null, sid) + v = vu + w * (vc - vu) + if rescale > 0: + sc = vc.std(dim=(1, 2), keepdim=True) + sv = v.std(dim=(1, 2), keepdim=True).clamp_min(1e-6) + v = rescale * (v * sc / sv) + (1 - rescale) * v + else: + v = vc + z = z + dt * v + return z + + +@torch.inference_mode() +def evaluate_gen_x(model, val_items, stats, sa, sr, cfg, w, rescales, epoch, demo_dir, n_pairs, n_demo): + """Realistic eval for the cross-attention generator (per-stem RESTORED context). + `rescales` is a list of CFG-rescale strengths to score in one pass (context is decoded once + and shared); the FIRST is primary (drives ckpt selection + demos). Reports a loudness ratio + (median decoded gen-RMS / true-RMS) โ€” the metric that actually tracks the ~2x overshoot.""" + model.eval(); dev = cfg.device; D = cfg.latent_dim + by_pair = defaultdict(list) + for it in val_items: + by_pair[(it["root"], it["pair"])].append(it) + pairs = list(by_pair.items()) + freq = defaultdict(int) + for _, its in pairs: + for it in its: freq[it["target"]] += 1 + pairs.sort(key=lambda kv: min(freq[it["target"]] for it in kv[1])) + pairs = pairs[:n_pairs] + + def dec(path): return E._decode(sa, _fitT(torch.load(path, map_location="cpu").float(), cfg.T)) + per = {r: defaultdict(lambda: defaultdict(list)) for r in rescales} + remix = {r: [] for r in rescales} + flat = {"clean": defaultdict(list), "gen": defaultdict(list)} # latent-FAD (primary rescale) + frmx = {"clean": [], "gen": []} + for pi, ((root, pid), its) in enumerate(pairs): + cache = {} + def dec_c(path): + if path not in cache: cache[path] = dec(path) + return cache[path] + for it in its: + s = it["target"]; sid = torch.tensor([STEM_ID[s]], device=dev) + src = "restored" if "restored" in it["ctx_by_src"] else next(iter(it["ctx_by_src"])) + cpaths = it["ctx_by_src"][src] + tm, ts = stats["tgt"][s] + ctx = torch.zeros(1, CTX_K, D, cfg.T); ids = torch.zeros(1, CTX_K, dtype=torch.long) + mask = torch.zeros(1, CTX_K, dtype=torch.bool) + for j, o in enumerate(sorted(cpaths)): + x = _fitT(torch.load(cpaths[o], map_location="cpu").float(), cfg.T) + sm, sd = stats["src"][o] + ctx[0, j] = (x - sm[:, None]) / sd[:, None]; ids[0, j] = STEM_ID[o]; mask[0, j] = True + ctx = ctx.to(dev); ids = ids.to(dev); mask = mask.to(dev) + at = dec_c(it["tgt_path"]); rms_t = at.pow(2).mean().sqrt().item() + 1e-9 + ctx_audio = sum(dec_c(p) for p in cpaths.values()) + clean_others = sum(dec_c(f"{root}/{pid}/latents/clean/{o}.pt") for o in cpaths) + for r in rescales: + g1 = sample_gen_x(model, ctx, ids, mask, sid, cfg.sample_steps, w, r) + g2 = sample_gen_x(model, ctx, ids, mask, sid, cfg.sample_steps, w, r) + per[r][s]["divers"].append((g1 - g2).flatten().pow(2).mean().sqrt().item()) + ag = E._decode(sa, _destd(g1[0], tm, ts)) + per[r][s]["stft_gen"].append(E.multi_stft(ag.mean(0), at.mean(0))) + per[r][s]["loud"].append(ag.pow(2).mean().sqrt().item() / rms_t) + remix[r].append(E.multi_stft((ctx_audio + ag).mean(0), (clean_others + at).mean(0))) + if r == rescales[0]: # latent-FAD (primary) + rg = _destd(g1[0], tm, ts).detach().cpu() + rc = _fitT(torch.load(it["tgt_path"], map_location="cpu").float(), cfg.T) + flat["gen"][s].append(rg); flat["clean"][s].append(rc) + ctx_oth = sum(_fitT(torch.load(cpaths[o], map_location="cpu").float(), cfg.T) for o in cpaths) + cln_oth = sum(_fitT(torch.load(f"{root}/{pid}/latents/clean/{o}.pt", map_location="cpu").float(), cfg.T) for o in cpaths) + frmx["gen"].append(ctx_oth + rg); frmx["clean"].append(cln_oth + rc) + if r == rescales[0] and pi < n_demo: + d = Path(demo_dir) / f"epoch_{epoch:04d}" / f"{pid}"; d.mkdir(parents=True, exist_ok=True) + import soundfile as sf + sf.write(d / f"{s}_gen.wav", ag.T.numpy(), sr); sf.write(d / f"{s}_true.wav", at.T.numpy(), sr) + sf.write(d / f"{s}_REMIXgen.wav", (ctx_audio + ag).T.numpy(), sr) + sf.write(d / f"{s}_REMIXcleanref.wav", (clean_others + at).T.numpy(), sr) + primary = {} + for r in rescales: + print(f"\n=== gen eval epoch {epoch} (pairs={len(pairs)}, CFG w={w} rescale={r}) ===") + print(f"{'stem':7} {'n':>4} | stft gen-vs-true | diversity | loud gen/true") + out = {} + for s in STEMS: + if not per[r][s]["stft_gen"]: continue + m = float(np.mean(per[r][s]["stft_gen"])); dv = float(np.mean(per[r][s]["divers"])) + ld = float(np.median(per[r][s]["loud"])) + out[s] = {"stft": m, "divers": dv, "loud": ld} + print(f"{s:7} {len(per[r][s]['stft_gen']):>4} | {m:.3f} | {dv:.3f} | {ld:.2f}x") + if remix[r]: + out["remix"] = float(np.mean(remix[r])) + print(f"{'REMIX':7} {len(remix[r]):>4} | {out['remix']:.3f} (others+gen vs others+true)") + if r == rescales[0]: primary = out + # latent-FAD: generated-stem distribution vs clean (the right one-to-many metric), per stem + REMIX + try: + print(f"{'stem':7} | latent-FAD gen-vs-clean (lower=closer)") + for s in STEMS: + if not flat["gen"][s]: continue + fo = fad.fad(fad.latent_set(flat["clean"][s]), fad.latent_set(flat["gen"][s])) + primary.setdefault(s, {})["fad"] = fo + print(f"{s:7} | {fo:.3f}") + if frmx["gen"]: + primary["fad_remix"] = fad.fad(fad.latent_set(frmx["clean"]), fad.latent_set(frmx["gen"])) + print(f"{'FAD-REMIX':12} | {primary['fad_remix']:.3f}") + except Exception as e: + print(f"[gen eval] latent-FAD skipped: {e}") + print(f"(demo -> {demo_dir}/epoch_{epoch:04d})") + return primary + + +# ---------------- flow: noise -> target, CFG ---------------- +def _endpoint_dist(z, tt, vhat, tgt, sid, dist_kind, dist_w): + """Cheap one-step endpoint distributional term (reuses vhat; no sampling). + Linear flow: x1_hat = z + (1-t)*vhat. Matches the GENERATED endpoint distribution to the + real one per stem -> directly targets one-to-many / perceptual quality (FAD-as-loss).""" + if dist_w <= 0: + return tgt.new_zeros(()) + x1 = z + (1 - tt) * vhat + return dist_w * DL.dist_term(dist_kind, x1, tgt, sid) + + +def fm_loss_gen(model, ctx, tgt, sid, cfg_drop, dist_kind="none", dist_w=0.0): + """Conditional flow matching from noise to target. CFG: drop context per-sample. + Optional endpoint distributional aux (dist_w>0).""" + B = tgt.shape[0] + if cfg_drop > 0: + m = (torch.rand(B, device=tgt.device) < cfg_drop).float()[:, None, None] + ctx = ctx * (1 - m) # zero context = unconditional token + x0 = torch.randn_like(tgt) + t = torch.rand(B, device=tgt.device); tt = t[:, None, None] + z = (1 - tt) * x0 + tt * tgt + v = tgt - x0 + vhat = model(z, t, ctx, sid) + return F.mse_loss(vhat, v) + _endpoint_dist(z, tt, vhat, tgt, sid, dist_kind, dist_w) + + +@torch.inference_mode() +def sample_gen(model, ctx, sid, steps, w, generator=None): + """Integrate noise->target with classifier-free guidance scale w.""" + z = torch.randn(ctx.shape, generator=generator, device=ctx.device) + zero = torch.zeros_like(ctx) + dt = 1.0 / steps + for i in range(steps): + t = torch.full((ctx.shape[0],), i * dt, device=ctx.device) + vc = model(z, t, ctx, sid) + if w != 1.0: + vu = model(z, t, zero, sid) + v = vu + w * (vc - vu) + else: + v = vc + z = z + dt * v + return z + + +# ---------------- eval: decode, remix, diversity ---------------- +def _destd(z, mu, sd): return z * sd[:, None].to(z.device) + mu[:, None].to(z.device) + + +@torch.inference_mode() +def evaluate_gen(model, val_items, stats, sa, sr, cfg, w, epoch, demo_dir, n_pairs, n_demo): + """End-to-end realistic eval: context = RESTORED others (as at inference); generate target; + remix = restored_others + generated vs the true CLEAN mix (clean_others + clean_target).""" + model.eval(); dev = cfg.device + by_pair = defaultdict(list) + for it in val_items: + by_pair[(it["root"], it["pair"])].append(it) + pairs = list(by_pair.items()) + freq = defaultdict(int) # stratify: rarest target stem first (bass coverage) + for _, its in pairs: + for it in its: freq[it["target"]] += 1 + pairs.sort(key=lambda kv: min(freq[it["target"]] for it in kv[1])) + pairs = pairs[:n_pairs] + + def dec(path): + return E._decode(sa, _fitT(torch.load(path, map_location="cpu").float(), cfg.T)) + + per = defaultdict(lambda: defaultdict(list)); remix = defaultdict(list) + flat = {"clean": defaultdict(list), "gen": defaultdict(list)} # latent-FAD: stem -> [ [256,T] ] + frmx = {"clean": [], "gen": []} # latent-FAD REMIX (per-pair latent sum) + for pi, ((root, pid), its) in enumerate(pairs): + cache = {} + def dec_c(path): + if path not in cache: cache[path] = dec(path) + return cache[path] + for it in its: + s = it["target"]; sid = torch.tensor([STEM_ID[s]], device=dev) + src = "restored" if "restored" in it["ctx_by_src"] else next(iter(it["ctx_by_src"])) + cpaths = it["ctx_by_src"][src] # {other_stem: restored path} + tm, ts = stats["tgt"][s]; cm, cs = stats["ctx"][s] + ctx_raw = sum(_fitT(torch.load(p, map_location="cpu").float(), cfg.T) for p in cpaths.values()) + ctx = ((ctx_raw - cm[:, None]) / cs[:, None]).to(dev)[None] + g1 = sample_gen(model, ctx, sid, cfg.sample_steps, w) + g2 = sample_gen(model, ctx, sid, cfg.sample_steps, w) # 2nd sample for diversity + per[s]["divers"].append((g1 - g2).flatten().pow(2).mean().sqrt().item()) + ag = E._decode(sa, _destd(g1[0], tm, ts)); at = dec_c(it["tgt_path"]) + per[s]["stft_gen"].append(E.multi_stft(ag.mean(0), at.mean(0))) + rg = _destd(g1[0], tm, ts).detach().cpu() # generated raw latent + rc = _fitT(torch.load(it["tgt_path"], map_location="cpu").float(), cfg.T) # clean target raw + flat["gen"][s].append(rg); flat["clean"][s].append(rc) + clean_oth_lat = sum(_fitT(torch.load(f"{root}/{pid}/latents/clean/{o}.pt", map_location="cpu").float(), cfg.T) for o in cpaths) + frmx["gen"].append(ctx_raw + rg); frmx["clean"].append(clean_oth_lat + rc) + ctx_audio = sum(dec_c(p) for p in cpaths.values()) # restored others + clean_others = sum(dec_c(f"{root}/{pid}/latents/clean/{o}.pt") for o in cpaths) # clean others + remix["stft_out"].append(E.multi_stft((ctx_audio + ag).mean(0), (clean_others + at).mean(0))) + if pi < n_demo: + d = Path(demo_dir) / f"epoch_{epoch:04d}" / f"{pid}"; d.mkdir(parents=True, exist_ok=True) + import soundfile as sf + sf.write(d / f"{s}_gen.wav", ag.T.numpy(), sr); sf.write(d / f"{s}_true.wav", at.T.numpy(), sr) + sf.write(d / f"{s}_REMIXgen.wav", (ctx_audio + ag).T.numpy(), sr) + sf.write(d / f"{s}_REMIXcleanref.wav", (clean_others + at).T.numpy(), sr) + print(f"\n=== gen eval epoch {epoch} (pairs={len(pairs)}, CFG w={w}) ===") + print(f"{'stem':7} {'n':>4} | stft gen-vs-true | sample diversity") + out = {} + for s in STEMS: + if not per[s]["stft_gen"]: continue + m = float(np.mean(per[s]["stft_gen"])); dv = float(np.mean(per[s]["divers"])) + out[s] = {"stft": m, "divers": dv} + print(f"{s:7} {len(per[s]['stft_gen']):>4} | {m:.3f} | {dv:.3f}") + if remix["stft_out"]: + out["remix"] = float(np.mean(remix["stft_out"])) + print(f"{'REMIX':7} {len(remix['stft_out']):>4} | {out['remix']:.3f} (others+gen vs others+true)") + # latent-FAD: distribution of GENERATED stems vs CLEAN (the right one-to-many metric; STFT/REMIX + # penalize valid-but-different samples, this does not). Per stem + REMIX, lower = closer to real. + try: + print(f"{'stem':7} | latent-FAD gen-vs-clean (lower=closer)") + for s in STEMS: + if not flat["gen"][s]: continue + fo = fad.fad(fad.latent_set(flat["clean"][s]), fad.latent_set(flat["gen"][s])) + out.setdefault(s, {})["fad"] = fo + print(f"{s:7} | {fo:.3f}") + if frmx["gen"]: + out["fad_remix"] = fad.fad(fad.latent_set(frmx["clean"]), fad.latent_set(frmx["gen"])) + print(f"{'FAD-REMIX':12} | {out['fad_remix']:.3f}") + except Exception as e: + print(f"[gen eval] latent-FAD skipped: {e}") + print(f"(demo -> {demo_dir}/epoch_{epoch:04d})") + return out + + +# ---------------- main ---------------- +def main(): + c = Cfg(); p = argparse.ArgumentParser() + p.add_argument("--cache-roots", default=",".join(c.cache_roots)) + p.add_argument("--val-cache-roots", default=",".join(c.val_cache_roots)) + p.add_argument("--out-dir", default=str(Path(c.out_dir).parent / "gen_v1")) + p.add_argument("--target-stems", default="bass,drums") + p.add_argument("--context-sources", default="restored,degraded", + help="context drawn (uniformly per item) from these; NEVER 'clean' at inference") + p.add_argument("--epochs", type=int, default=200) + p.add_argument("--batch", type=int, default=256) + p.add_argument("--lr", type=float, default=1e-3) + p.add_argument("--hidden", type=int, default=512) + p.add_argument("--depth", type=int, default=6) + p.add_argument("--arch", default="conv", choices=["conv", "attn", "xattn"], + help="conv=FiLM-conv CondFlow (v2); attn=Conformer-lite self-attn velocity net; " + "xattn=cross-attention DiT (per-stem context tokens, not summed)") + p.add_argument("--heads", type=int, default=8) + p.add_argument("--sigma", type=float, default=1.0) # unused (noise-start); kept for parity + p.add_argument("--cfg-drop", type=float, default=0.1) + p.add_argument("--cfg-w", type=float, default=2.0) + p.add_argument("--cfg-rescale", type=float, default=0.0, + help="xattn only: norm-preserving CFG strength (0=off, ~0.7=Sony) โ€” loudness fix") + p.add_argument("--dist-loss", default="none", choices=["none", "var", "mmd", "moment"], + help="cheap latent distributional aux via one-step endpoint (tackles one-to-many)") + p.add_argument("--dist-loss-w", type=float, default=0.0) + p.add_argument("--sample-steps", type=int, default=40) + p.add_argument("--eval-every", type=int, default=25) + p.add_argument("--eval-pairs", type=int, default=150) + p.add_argument("--demo-pairs", type=int, default=8) + p.add_argument("--device", default="cuda") + p.add_argument("--num-workers", type=int, default=4) + p.add_argument("--smoke", action="store_true") + a = p.parse_args() + if a.smoke: + a.epochs, a.eval_every, a.eval_pairs, a.demo_pairs, a.num_workers = 2, 1, 6, 2, 0 + + random.seed(0); torch.manual_seed(0) + out = Path(a.out_dir); out.mkdir(parents=True, exist_ok=True) + json.dump(vars(a), open(out / "config.json", "w"), indent=2) + tgt_stems = [s for s in a.target_stems.split(",") if s] + cfg = type(c)(**{**c.__dict__, "device": a.device, "sample_steps": a.sample_steps, "T": c.T}) + + tr_roots = [x for x in a.cache_roots.split(",") if x] + va_roots = [x for x in a.val_cache_roots.split(",") if x] or tr_roots + sources = tuple(x for x in a.context_sources.split(",") if x) + print(f"[gen] context sources={sources} (target=clean)") + tr_items = index_clips(tr_roots, tgt_stems, sources) + va_items = index_clips(va_roots, tgt_stems, sources) + if a.smoke: + tr_items, va_items = tr_items[:1500], va_items[:200] + from collections import Counter + print(f"[gen] train={len(tr_items)} {dict(Counter(i['target'] for i in tr_items))}") + print(f"[gen] val ={len(va_items)} {dict(Counter(i['target'] for i in va_items))}") + + is_x = a.arch == "xattn" + stats_f = out / ("gen_stats_x.pt" if is_x else "gen_stats.pt") + if stats_f.exists(): + stats = torch.load(stats_f) + else: + stats = (build_gen_stats_x if is_x else build_gen_stats)(tr_items, cfg.T, cfg.latent_dim) + torch.save(stats, stats_f) + ds = GenDataX(tr_items, stats, cfg.T, cfg.latent_dim) if is_x else GenData(tr_items, stats, cfg.T) + loader = DataLoader(ds, batch_size=a.batch, shuffle=True, drop_last=True, + num_workers=a.num_workers, pin_memory=True) + if is_x: + model = XAttnCondFlow(cfg.latent_dim, a.hidden, a.depth, n_stems=len(STEMS), + stem_emb=cfg.stem_emb_dim, heads=a.heads).to(a.device) + elif a.arch == "attn": + model = AttnCondFlow(cfg.latent_dim, a.hidden, a.depth, n_stems=len(STEMS), + stem_emb=cfg.stem_emb_dim, heads=a.heads).to(a.device) + else: + model = CondFlow(cfg.latent_dim, a.hidden, a.depth, n_stems=len(STEMS), + stem_emb=cfg.stem_emb_dim).to(a.device) + print(f"[gen] {a.arch} params={model.num_params()/1e6:.2f}M cfg_drop={a.cfg_drop} cfg_w={a.cfg_w}") + opt = torch.optim.AdamW(model.parameters(), a.lr, weight_decay=1e-4) + sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, a.epochs * max(1, len(loader))) + sa, sr = E.load_same(cfg) + + best = float("inf") + for ep in range(a.epochs): + model.train(); t0 = time.time(); tot = 0.0 + for b in loader: + tgt = b["tgt"].to(a.device); sid = b["stem_id"].to(a.device) + if is_x: + loss = fm_loss_gen_x(model, b["ctx_stems"].to(a.device), b["ctx_ids"].to(a.device), + b["ctx_mask"].to(a.device), tgt, sid, a.cfg_drop, + a.dist_loss, a.dist_loss_w) + else: + loss = fm_loss_gen(model, b["ctx"].to(a.device), tgt, sid, a.cfg_drop, + a.dist_loss, a.dist_loss_w) + opt.zero_grad(); loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # flow-matching can spike on small data + opt.step(); sched.step(); tot += loss.item() + print(f"epoch {ep+1}/{a.epochs} loss={tot/max(1,len(loader)):.4f} ({time.time()-t0:.0f}s)") + if (ep + 1) % a.eval_every == 0 or ep == a.epochs - 1: + if is_x: + rescales = [a.cfg_rescale] + ([0.0] if a.cfg_rescale != 0.0 else []) + m = evaluate_gen_x(model, va_items, stats, sa, sr, cfg, a.cfg_w, rescales, + ep + 1, str(out / "demo"), a.eval_pairs, a.demo_pairs) + else: + m = evaluate_gen(model, va_items, stats, sa, sr, cfg, a.cfg_w, ep + 1, str(out / "demo"), + a.eval_pairs, a.demo_pairs) + ck = {"model": model.state_dict(), "stats": stats, "args": vars(a), "epoch": ep + 1, "metrics": m} + torch.save(ck, out / "ckpt.pt") + r = float(m.get("remix", float("inf"))) + if r < best: + best = r; torch.save(ck, out / "ckpt_best.pt"); print(f" ** best REMIX={r:.3f} @ ep{ep+1}") + print(f"done. best REMIX={best:.3f}") + + +if __name__ == "__main__": + main() diff --git a/restoflow/gen_fad_probe.py b/restoflow/gen_fad_probe.py new file mode 100644 index 0000000000000000000000000000000000000000..bbb09d50629dc8ca98534328ccb71e03dcbfad5f --- /dev/null +++ b/restoflow/gen_fad_probe.py @@ -0,0 +1,62 @@ +"""Apples-to-apples generator FAD: load saved gen runs and re-run the (now FAD-enabled) eval once +each, on identical val pairs, so per-stem + REMIX latent-FAD is comparable across the capacity ladder +(conv-10M vs xattn-40M vs xattn-63M). Reuses evaluate_gen / evaluate_gen_x. + +Run: python -m restoflow.gen_fad_probe --runs gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M --device cuda +""" +from __future__ import annotations +import argparse +from pathlib import Path +import torch +from .config import Cfg, STEMS +from .model import CondFlow, AttnCondFlow, XAttnCondFlow +from . import eval as E +from .gen import index_clips, evaluate_gen, evaluate_gen_x + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--runs", default="gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M") + p.add_argument("--val-cache-roots", default=str(BASE / "demucs_results_val_full")) + p.add_argument("--device", default="cuda") + p.add_argument("--eval-pairs", type=int, default=120) + a = p.parse_args() + dev = a.device + cfg = Cfg(device=dev, T=32) + sa, sr = E.load_same(cfg) + va_roots = [x for x in a.val_cache_roots.split(",") if x] + + for run in [r for r in a.runs.split(",") if r]: + ck_path = BASE / "restoflow_runs" / run / "ckpt_best.pt" + if not ck_path.exists(): + print(f"\n##### {run}: no ckpt_best.pt โ€” skip"); continue + ck = torch.load(ck_path, map_location="cpu"); ga = ck["args"] + arch = ga.get("arch", "conv") + if arch == "xattn": + m = XAttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), + stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8)) + elif arch == "attn": + m = AttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), + stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8)) + else: + m = CondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim) + m.load_state_dict(ck["model"]); m.eval().to(dev) + stats = ck["stats"] + tgt_stems = [s for s in ga.get("target_stems", "bass").split(",") if s] + sources = tuple(x for x in ga.get("context_sources", "restored,degraded").split(",") if x) + c = type(cfg)(**{**cfg.__dict__, "sample_steps": ga.get("sample_steps", 40)}) + va = index_clips(va_roots, tgt_stems, sources) + print(f"\n##### {run} arch={arch} {ga['hidden']}x{ga['depth']} " + f"params={sum(p_.numel() for p_ in m.parameters())/1e6:.1f}M val={len(va)} #####") + demo = str(BASE / "restoflow_runs" / run / "_fadprobe_demo") + if arch == "xattn": + evaluate_gen_x(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), [ga.get("cfg_rescale", 0.0)], + 0, demo, a.eval_pairs, 0) + else: + evaluate_gen(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), 0, demo, a.eval_pairs, 0) + + +if __name__ == "__main__": + main() diff --git a/restoflow/lit_eval.py b/restoflow/lit_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..088c7b39b6bc234014532d5a44d68933c3aed55b --- /dev/null +++ b/restoflow/lit_eval.py @@ -0,0 +1,239 @@ +"""Literature-comparable evaluation: place our restoration on the SAME axes the BWE / music-restoration +papers report, so we can see where we stand. + +Shared baselines & what they report: + - AERO (ICASSP'23, spectral audio super-res): LSD, ViSQOL, MUSHRA โ€” VCTK / MUSDB18. + - BigWavGAN ('23, wave GAN music super-res): LSD, SI-SDR, ViSQOL โ€” MUSDB18. + - BABE-2 / Diffusion Generative Equalizer (DAFx'24, Moliner): FAD, LSD โ€” music restoration. + +Shared metrics we compute here on OUR val set, for the BWE-comparable unit = the mix of present +(restoration-class) stems: degraded (lower bound) and OURS (advramp restoration), each vs clean (oracle): + - LSD log-spectral distance (lower=better; the BWE standard) + - LSD-HF LSD restricted to >hf_hz (the high band โ€” our brightness/hiss question, where BWE lives) + - SI-SDR scale-invariant SDR (higher=better; fidelity) + - FAD Frechet Audio Distance (VGGish + CLAP-music; lower=better; the generative-restoration metric) + +CAVEAT printed in the report: our dataset + degradation differ from each paper's, so absolute numbers are +NOT a head-to-head ranking โ€” they place us in the same REGIME/units. For a true head-to-head, re-run on +MUSDB18-HQ with a matched low-pass degradation (--musdb path, future). + +Run: python -m restoflow.lit_eval --device cuda --n-pairs 120 +""" +from __future__ import annotations +import argparse, csv, math, random +from collections import defaultdict +from pathlib import Path +import numpy as np +import torch + +from .config import Cfg, STEMS +from . import eval as E + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") +EPS = 1e-8 + + +def _stft_mag(x, n_fft=2048, hop=512): + w = torch.hann_window(n_fft) + X = torch.stft(torch.from_numpy(x).float(), n_fft, hop_length=hop, window=w, return_complex=True) + return X.abs().numpy() # [F, T] + + +def lsd(ref, est, sr=44100, n_fft=2048, hop=512, hf_hz=None): + """Log-spectral distance (dB-domain power), mono. hf_hz -> restrict to the high band only.""" + r = _stft_mag(ref.mean(0), n_fft, hop); e = _stft_mag(est.mean(0), n_fft, hop) + T = min(r.shape[1], e.shape[1]); r, e = r[:, :T], e[:, :T] + lr = np.log10(r ** 2 + EPS); le = np.log10(e ** 2 + EPS) + if hf_hz is not None: + f0 = int(round(hf_hz / (sr / 2) * (r.shape[0] - 1))) + lr, le = lr[f0:], le[f0:] + return float(np.mean(np.sqrt(np.mean((lr - le) ** 2, axis=0)))) + + +def si_sdr(ref, est): + """Scale-invariant SDR (dB), mono, length-aligned.""" + r = ref.mean(0).astype(np.float64); e = est.mean(0).astype(np.float64) + n = min(len(r), len(e)); r, e = r[:n], e[:n] + r = r - r.mean(); e = e - e.mean() + alpha = np.dot(e, r) / (np.dot(r, r) + EPS) + s = alpha * r; noise = e - s + return float(10 * np.log10((np.dot(s, s) + EPS) / (np.dot(noise, noise) + EPS))) + + +def _best_lag(r, e, max_lag): + """FFT cross-correlation lag of est vs ref within +-max_lag (no scipy dep).""" + n = 1 << int(math.ceil(math.log2(2 * len(r) + 1))) + cc = np.fft.irfft(np.fft.rfft(e, n) * np.conj(np.fft.rfft(r, n)), n) + cc = np.concatenate([cc[-max_lag:], cc[:max_lag + 1]]) + return int(np.arange(-max_lag, max_lag + 1)[int(np.argmax(cc))]) + + +def align(ref, est, sr=44100, max_lag_s=0.05): + """Paper-faithful: correct ONLY the constant codec/window time-latency (sub-50ms cross-correlation) + so LSD/SI-SDR measure spectral/waveform quality, not a systematic offset. NO level-norm โ€” LSD should + count level differences, as the SR papers do (their models output level-matched audio).""" + n = min(ref.shape[1], est.shape[1]); ref, est = ref[:, :n], est[:, :n] + lag = _best_lag(ref.mean(0).astype(np.float64), est.mean(0).astype(np.float64), int(max_lag_s * sr)) + if lag > 0: est = est[:, lag:]; ref = ref[:, :est.shape[1]] + elif lag < 0: ref = ref[:, -lag:]; est = est[:, :ref.shape[1]] + L = min(ref.shape[1], est.shape[1]); return ref[:, :L], est[:, :L] + + +def add_metrics(M, tag, ref, est, sr, hf_hz): + ref, est = align(ref, est, sr) + M[tag]["LSD"].append(lsd(ref, est, sr)); M[tag]["LSD-HF"].append(lsd(ref, est, sr, hf_hz=hf_hz)) + M[tag]["SI-SDR"].append(si_sdr(ref, est)) + + +def codec_passthrough(audio, APP): + """SAME encode->decode of a mix (3s windows, no separation/restoration) = the autoencoder codec FLOOR + our system cannot beat. Lets MUSDB show how much of the gap is the codec vs the restoration.""" + SR = APP.SR; chunk = SR * 3; outs = [] + for i in range(0, audio.shape[1], chunk): + seg = audio[:, i:i + chunk]; o = seg.shape[1] + if o < chunk: seg = np.pad(seg, ((0, 0), (0, chunk - o))) + outs.append(np.asarray(APP.decode(APP.encode(seg)))) + return np.concatenate(outs, 1) if outs else audio + + +def _fad_rows(aud, sr): + """FAD-VGGish + FAD-CLAP-music for {deg,ours} vs clean. Returns list of (name, deg, ours) or notes.""" + rows = [] + try: + from . import fad as FAD + except Exception as e: + return [("FAD", f"skipped ({e})", "")] + to_t = lambda L: [torch.from_numpy(np.ascontiguousarray(x)).float() for x in L] # embed_set wants tensors + for name, klass in (("VGGish", "VGGish"), ("CLAP", "CLAP")): + try: + emb = getattr(FAD, klass)() + ec, ed, eo = (emb.embed_set(to_t(aud[k]), sr) for k in ("clean", "deg", "ours")) + D = ec.shape[1]; Nmin = min(len(ec), len(ed), len(eo)) + valid = Nmin >= 2 * D # need embeddings >> dim for a full-rank covariance + sh = not valid # empirical cov when valid; shrink ONLY (flagged) when under-sampled + tag = f"FAD-{name} (N={Nmin},D={D}{' โœ“' if valid else ' โš shrink'})" + rows.append((tag, f"{FAD.fad(ec, ed, shrink=sh):.3f}", f"{FAD.fad(ec, eo, shrink=sh):.3f}")) + except Exception as e: + rows.append((f"FAD-{name}", f"skip:{type(e).__name__}", "")) + return rows + + +def _print_table(label, M, aud, sr, n, tags=("deg", "ours")): + titles = {"deg": "degraded", "ours": "OURS", "codecfloor": "codec-floor"} + print(f"\n=== {label} ยท {n} items (vs clean oracle) ===") + print(f"{'metric':12} " + " ".join(f"{titles.get(t, t):>12}" for t in tags)) + for m in ("LSD", "LSD-HF", "SI-SDR"): + arrow = "โ†‘" if m == "SI-SDR" else "โ†“" + print(f"{m:12} " + " ".join(f"{np.mean(M[t][m]):>12.3f}" for t in tags) + f" ({arrow} better)") + for name, dv, ov in _fad_rows(aud, sr): # FAD only for deg/ours vs clean set + print(f" {name:34} deg={dv:>9} ours={ov:>9}") + + +def run_musdb(a): + """Full-pipeline eval on MUSDB18-HQ: clean mixture -> deterministic low-pass degrade -> our app + pipeline (separate->restore->remix) -> metrics vs clean. Same axes as the BWE papers' MUSDB tables.""" + import os, glob, tempfile, soundfile as sf, librosa + os.environ.setdefault("RESTOFLOW_DEVICE", a.device) + import restoflow.app as APP + APP.models() + SR = APP.SR; exc = int(a.excerpt_s * SR) + M = {t: defaultdict(list) for t in ("deg", "ours", "codecfloor")}; aud = {"clean": [], "deg": [], "ours": []} + tracks = sorted([d for d in glob.glob(str(Path(a.musdb_root) / "*")) if Path(d).is_dir()]) + random.shuffle(tracks); tracks = tracks[:a.musdb_n]; used = 0 + for td in tracks: + mixp = Path(td) / "mixture.wav" + if mixp.exists(): + clean, _ = sf.read(mixp, dtype="float32"); clean = clean.T + else: + sts = [sf.read(Path(td) / f"{s}.wav", dtype="float32")[0].T for s in STEMS if (Path(td) / f"{s}.wav").exists()] + if not sts: continue + clean = sum(sts) + if clean.ndim == 1: clean = np.stack([clean, clean]) + if clean.shape[1] < exc: continue + st = (clean.shape[1] - exc) // 2; clean = np.ascontiguousarray(clean[:, st:st + exc]) + deg = np.stack([librosa.resample(clean[c], orig_sr=SR, target_sr=a.lp_sr) for c in range(2)]) # low-pass + deg = np.ascontiguousarray(np.stack([librosa.resample(deg[c], orig_sr=a.lp_sr, target_sr=SR) + for c in range(2)])[:, :clean.shape[1]].astype(np.float32)) + if not np.isfinite(deg).all() or float(np.sqrt(np.mean(deg ** 2))) < 1e-5: + continue # skip silent/invalid excerpts + tmp = tempfile.mktemp(suffix=".wav"); sf.write(tmp, deg.T, SR, subtype="FLOAT") + out = APP.run_upload(tmp) + if not out or out[1] is None: # run_upload bailed (sep failed) + continue + ours = np.asarray(out[1][1]).T + codec = codec_passthrough(clean, APP) # SAME codec floor (no sep/restore) + add_metrics(M, "deg", clean, deg, SR, a.hf_hz) + add_metrics(M, "ours", clean, ours, SR, a.hf_hz) + add_metrics(M, "codecfloor", clean, codec, SR, a.hf_hz) + aud["clean"].append(clean); aud["deg"].append(deg); aud["ours"].append(ours); used += 1 + _print_table(f"MUSDB18-HQ (raw-clean ref; low-pass {a.lp_sr}Hz; full pipeline incl. separation+codec)", + M, aud, SR, used, tags=("deg", "ours", "codecfloor")) + print(" (codec-floor = SAME encode->decode of clean, NO separation/restoration = the ceiling our latent") + print(" pipeline can reach; the ours->codecfloor gap is restoration+separation, codecfloor->0 is the codec tax.)") + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--device", default="cuda"); p.add_argument("--n-pairs", type=int, default=120) + p.add_argument("--hf-hz", type=float, default=4000.0, help="LSD-HF band cutoff") + p.add_argument("--val-root", default=str(BASE / "demucs_results_val_full")) + p.add_argument("--musdb-root", default="/media/maindisk/melkor169/drums_dereverb/data/gmd_musdb18hq_stereo") + p.add_argument("--musdb-n", type=int, default=25); p.add_argument("--excerpt-s", type=float, default=9.0) + p.add_argument("--lp-sr", type=int, default=8000, help="low-pass cutoff = lp_sr/2 (BWE degradation)") + p.add_argument("--skip-ourval", action="store_true"); p.add_argument("--skip-musdb", action="store_true") + a = p.parse_args(); dev = a.device; random.seed(0) + cfg = Cfg(device=dev, T=32) + sa, sr = E.load_same(cfg) + root = Path(a.val_root) + + # ---- (1) OUR val set (cached latents; restoration-class stem mix) ---- + if not a.skip_ourval: + rows = list(csv.DictReader(open(root / "metadata.csv"))) + by_pair = defaultdict(list) + for r in rows: + by_pair[r["pair_id"]].append(r) + pairs = list(by_pair.items()); random.shuffle(pairs); pairs = pairs[:a.n_pairs] + M = {"deg": defaultdict(list), "ours": defaultdict(list)}; aud = {"clean": [], "deg": [], "ours": []}; used = 0 + with torch.inference_mode(): + for pid, srows in pairs: + cl = dg = rr = None + for r in srows: + s = r.get("stem") + if s not in STEMS or r.get("class") != "restoration": + continue + cp = r.get("clean_latent_path"); dp = r.get("deg_latent_path") + rp = root / pid / "latents" / "restored" / f"{s}.pt" + if not (cp and dp and (root / cp).exists() and (root / dp).exists() and rp.exists()): + continue + c = E._decode(sa, torch.load(root / cp, map_location="cpu").float()) + d = E._decode(sa, torch.load(root / dp, map_location="cpu").float()) + o = E._decode(sa, torch.load(rp, map_location="cpu").float()) + cl = c if cl is None else cl[:, :c.shape[1]] + c[:, :cl.shape[1]] + dg = d if dg is None else dg[:, :d.shape[1]] + d[:, :dg.shape[1]] + rr = o if rr is None else rr[:, :o.shape[1]] + o[:, :rr.shape[1]] + if cl is None: + continue + cl, dg, rr = (x.numpy() if hasattr(x, "numpy") else x for x in (cl, dg, rr)) + n = min(cl.shape[1], dg.shape[1], rr.shape[1]); cl, dg, rr = cl[:, :n], dg[:, :n], rr[:, :n] + for tag, est in (("deg", dg), ("ours", rr)): + add_metrics(M, tag, cl, est, sr, a.hf_hz) # aligned + level-matched + aud["clean"].append(cl); aud["deg"].append(dg); aud["ours"].append(rr); used += 1 + # NOTE: here clean/deg/ours are ALL SAME-decoded -> same codec domain -> the delta is fair; the + # absolute LSD is vs codec-clean (not raw), so it isn't directly the papers' axis (they use raw clean). + _print_table("OUR val (codec-domain ref; restoration delta is fair, absolute != papers)", M, aud, sr, used) + + # ---- (2) MUSDB18-HQ (shared dataset; low-pass degrade; full pipeline) ---- + if not a.skip_musdb and Path(a.musdb_root).exists(): + run_musdb(a) + + print("\nCurrent (2024-2026) baselines on this task family + what they report:") + print(" AudioSR (2024) latent-diffusion versatile SR->48k | LSD,LSD-HF,ViSQOL,FAD | THE std baseline; ALSO latent => shares our codec floor") + print(" FlashSR (2025) 1-step diffusion-distill, SOTA music SR | LSD,ViSQOL,FAD | per-cutoff music SR") + print(" UniverSR (2025) vocoder-free FLOW-MATCHING SR, SOTA | LSD,LSD-HF,2f-model | 8/12/16/24->48k (no codec => no floor)") + print(" BABE-2 (2024) BLIND generative music restoration | FAD(CLAP),LSD | closest to OUR blind/unknown-degradation setting") + print(" CAVEAT: those are KNOWN-downsample, full-mix, single-stage SR; OURS is blind + per-stem + invents") + print(" missing instruments => no exact head-to-head. Shared axes = LSD / LSD-HF / FAD. (ViSQOL/2f-model: TODO.)") + + +if __name__ == "__main__": + main() diff --git a/restoflow/model.py b/restoflow/model.py new file mode 100644 index 0000000000000000000000000000000000000000..ecbbe452414921c6a4c91bd81bdb8a2a8f9d67e5 --- /dev/null +++ b/restoflow/model.py @@ -0,0 +1,338 @@ +"""Conditional flow velocity network over the SAME latent sequence [B,256,T]. + +Predicts velocity v(z_t, t | degraded, stem). Small (latents are tiny). One SHARED +net conditioned on stem id (FiLM) โ€” validated >= separate per-stem networks. +""" +from __future__ import annotations +import math +import torch +import torch.nn as nn + + +def sinusoidal(t: torch.Tensor, dim: int) -> torch.Tensor: # t [B] in [0,1] + half = dim // 2 + freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / max(half - 1, 1)) + a = t[:, None] * freqs[None] * 2 * math.pi + return torch.cat([a.sin(), a.cos()], dim=-1) + + +class FiLMBlock(nn.Module): + def __init__(self, h, cond_dim, k=5): + super().__init__() + self.conv1 = nn.Conv1d(h, h, k, padding=k // 2) + self.conv2 = nn.Conv1d(h, h, 1) + self.film = nn.Linear(cond_dim, 2 * h) + self.act = nn.GELU() + + def forward(self, x, cond): + g, b = self.film(cond).chunk(2, dim=-1) # [B,h] each + h = self.conv1(x) + h = h * (1 + g[:, :, None]) + b[:, :, None] + h = self.act(h) + h = self.conv2(h) + return x + h + + +class CondFlow(nn.Module): + def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, t_dim=64): + super().__init__() + self.stem_emb = nn.Embedding(n_stems, stem_emb) + cond_dim = t_dim + stem_emb + self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim)) + self.t_dim = t_dim + self.inp = nn.Conv1d(2 * d, hidden, 1) # concat(z_t, degraded) + self.blocks = nn.ModuleList([FiLMBlock(hidden, cond_dim) for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) # start near identity velocity + + def forward(self, z_t, t, deg, stem_id): + # z_t,deg [B,256,T]; t [B]; stem_id [B] + cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1) + h = self.inp(torch.cat([z_t, deg], dim=1)) + for blk in self.blocks: + h = blk(h, cond) + return self.out(h) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +class AttnCondFlow(nn.Module): + """Attention velocity net for the generator: same Conformer-lite block that won for the + restorer, conditioned on (t, stem) via FiLM. Global self-attention lets the invented stem + see the whole context window (conv-only CondFlow could not), which lands flow endpoints + closer to the clean manifold (the off-manifold endpoints are what decode too loud).""" + def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64, + heads=8, t_max=64): + super().__init__() + self.stem_emb = nn.Embedding(n_stems, stem_emb) + self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim)) + self.t_dim = t_dim + cond_dim = t_dim + stem_emb + self.inp = nn.Conv1d(2 * d, hidden, 1) # concat(z_t, ctx) + self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False) + self.blocks = nn.ModuleList([ConformerBlock(hidden, cond_dim, heads) for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) + + def forward(self, z_t, t, deg, stem_id): + cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1) + h = self.inp(torch.cat([z_t, deg], dim=1)) + self.pos[:, :, :z_t.shape[2]] + for blk in self.blocks: + h = blk(h, cond) + return self.out(h) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +class ConformerBlock(nn.Module): + """Pre-norm Conformer-lite: half-FFN -> MHSA -> FiLM depthwise-conv -> half-FFN. + Global self-attention (T~=32 is tiny) fixes the conv-only receptive-field gap; + per-block FiLM injects stem id at every depth instead of once at the input.""" + def __init__(self, h, cond_dim, heads=4, k=7, ff_mult=2): + super().__init__() + self.ln_ff1 = nn.LayerNorm(h) + self.ff1 = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h)) + self.ln_attn = nn.LayerNorm(h) + self.attn = nn.MultiheadAttention(h, heads, batch_first=True) + self.ln_conv = nn.LayerNorm(h) + self.pw1 = nn.Conv1d(h, 2 * h, 1) + self.dw = nn.Conv1d(h, h, k, padding=k // 2, groups=h) + self.gn = nn.GroupNorm(1, h) + self.film = nn.Linear(cond_dim, 2 * h) + self.pw2 = nn.Conv1d(h, h, 1) + self.act = nn.GELU() + self.ln_ff2 = nn.LayerNorm(h) + self.ff2 = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h)) + + def forward(self, x, cond): # x [B,h,T], cond [B,cond_dim] + xt = x.transpose(1, 2) # [B,T,h] + xt = xt + 0.5 * self.ff1(self.ln_ff1(xt)) + a = self.ln_attn(xt) + a, _ = self.attn(a, a, a, need_weights=False) + xt = xt + a + # conv module (channels-first) + c = self.ln_conv(xt).transpose(1, 2) # [B,h,T] + c = nn.functional.glu(self.pw1(c), dim=1) + c = self.gn(self.dw(c)) + g, b = self.film(cond).chunk(2, dim=-1) + c = self.act(c * (1 + g[:, :, None]) + b[:, :, None]) + c = self.pw2(c) + xt = xt + c.transpose(1, 2) + xt = xt + 0.5 * self.ff2(self.ln_ff2(xt)) + return xt.transpose(1, 2) + + +class CrossAttnBlock(nn.Module): + """DiT-style block for the generator: self-attn over the noised target frames -> + CROSS-attn into per-stem context tokens (keeps which instrument is which, vs the + old summed context) -> FiLM(t,stem) depthwise-conv -> FFN. A key-padding mask hides + absent context stems (and, on CFG drop, ALL real context -> only a learned null token).""" + def __init__(self, h, cond_dim, heads=8, k=7, ff_mult=2): + super().__init__() + self.ln_sa = nn.LayerNorm(h) + self.self_attn = nn.MultiheadAttention(h, heads, batch_first=True) + self.ln_ca = nn.LayerNorm(h) + self.cross_attn = nn.MultiheadAttention(h, heads, batch_first=True) + self.ln_conv = nn.LayerNorm(h) + self.pw1 = nn.Conv1d(h, 2 * h, 1) + self.dw = nn.Conv1d(h, h, k, padding=k // 2, groups=h) + self.gn = nn.GroupNorm(1, h) + self.film = nn.Linear(cond_dim, 2 * h) + self.pw2 = nn.Conv1d(h, h, 1) + self.act = nn.GELU() + self.ln_ff = nn.LayerNorm(h) + self.ff = nn.Sequential(nn.Linear(h, ff_mult * h), nn.GELU(), nn.Linear(ff_mult * h, h)) + + def forward(self, x, ctx_tok, cond, ctx_pad): # x [B,T,h], ctx_tok [B,M,h], cond [B,cond_dim], ctx_pad [B,M] (True=ignore) + a = self.ln_sa(x) + a, _ = self.self_attn(a, a, a, need_weights=False) + x = x + a + q = self.ln_ca(x) + c, _ = self.cross_attn(q, ctx_tok, ctx_tok, key_padding_mask=ctx_pad, need_weights=False) + x = x + c + cc = self.ln_conv(x).transpose(1, 2) # [B,h,T] + cc = nn.functional.glu(self.pw1(cc), dim=1) + cc = self.gn(self.dw(cc)) + g, b = self.film(cond).chunk(2, dim=-1) + cc = self.act(cc * (1 + g[:, :, None]) + b[:, :, None]) + cc = self.pw2(cc) + x = x + cc.transpose(1, 2) + x = x + self.ff(self.ln_ff(x)) + return x + + +class XAttnCondFlow(nn.Module): + """Cross-attention (DiT-style) conditional flow generator. Each CONTEXT stem latent is + encoded as its own token sequence (projected + per-instrument embedding + positional); + the noised target cross-attends into the union of those tokens. This replaces summing the + context (which discarded instrument identity) and is the principled 'transformer encoder' + conditioning from the bass-accompaniment literature (Sony arXiv:2402.01412). A learned + null token gives a well-defined unconditional pass for classifier-free guidance.""" + def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64, + heads=8, t_max=64): + super().__init__() + self.stem_emb = nn.Embedding(n_stems, stem_emb) + self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim)) + self.t_dim = t_dim + cond_dim = t_dim + stem_emb + self.inp = nn.Conv1d(d, hidden, 1) # noised target only (context via cross-attn) + self.ctx_proj = nn.Conv1d(d, hidden, 1) # shared per-stem context projection + self.ctx_stem_emb = nn.Embedding(n_stems, hidden) # which instrument each context token is + self.null_ctx = nn.Parameter(torch.randn(1, 1, hidden) * 0.02) + self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False) + self.blocks = nn.ModuleList([CrossAttnBlock(hidden, cond_dim, heads) for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) + + def _encode_ctx(self, ctx_stems, ctx_ids, ctx_mask): + # ctx_stems [B,K,d,T]; ctx_ids [B,K]; ctx_mask [B,K] (True=present) + B, K, d, T = ctx_stems.shape + x = self.ctx_proj(ctx_stems.reshape(B * K, d, T)) # [B*K,h,T] + x = x + self.pos[:, :, :T] + x = x.transpose(1, 2) # [B*K,T,h] + x = x + self.ctx_stem_emb(ctx_ids.reshape(B * K))[:, None, :] + tok = x.reshape(B, K * T, -1) + pad = (~ctx_mask)[:, :, None].expand(B, K, T).reshape(B, K * T) # True = ignore + null = self.null_ctx.expand(B, -1, -1) # [B,1,h] always attended + tok = torch.cat([null, tok], dim=1) + nullpad = torch.zeros(B, 1, dtype=torch.bool, device=tok.device) + return tok, torch.cat([nullpad, pad], dim=1) + + def forward(self, z_t, t, ctx_stems, ctx_ids, ctx_mask, stem_id): + cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1) + tok, pad = self._encode_ctx(ctx_stems, ctx_ids, ctx_mask) + x = self.inp(z_t) + self.pos[:, :, :z_t.shape[2]] # [B,h,T] + x = x.transpose(1, 2) # [B,T,h] + for blk in self.blocks: + x = blk(x, tok, cond, pad) + return self.out(x.transpose(1, 2)) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +class MixAttnCondFlow(nn.Module): + """STOCHASTIC restorer: attention velocity net for the deg-anchored flow bridge, WITH mix + conditioning. Combines the three validated levers โ€” attention (won deterministically), + mix-conditioning (config notes: helps drums most), and stochasticity (the literature fix + for the smeared transients that L2 regression averages away on percussive stems). The + deterministic restorers can't add transient detail (conditional-mean); this can sample it. + Signature model(z_t, t, deg, stem_id, mix) matches flow.fm_loss_mix / flow.sample_mix.""" + def __init__(self, d=256, hidden=512, depth=8, n_stems=4, stem_emb=32, t_dim=64, + heads=8, t_max=64, use_mix=True): + super().__init__() + self.use_mix = use_mix + self.stem_emb = nn.Embedding(n_stems, stem_emb) + self.t_mlp = nn.Sequential(nn.Linear(t_dim, t_dim), nn.GELU(), nn.Linear(t_dim, t_dim)) + self.t_dim = t_dim + cond_dim = t_dim + stem_emb + self.inp = nn.Conv1d(d * (3 if use_mix else 2), hidden, 1) # concat(z_t, deg, [mix]) + self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False) + self.blocks = nn.ModuleList([ConformerBlock(hidden, cond_dim, heads) for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) + + def forward(self, z_t, t, deg, stem_id, mix=None): + cond = torch.cat([self.t_mlp(sinusoidal(t, self.t_dim)), self.stem_emb(stem_id)], dim=-1) + parts = [z_t, deg, mix] if (self.use_mix and mix is not None) else [z_t, deg] + h = self.inp(torch.cat(parts, dim=1)) + self.pos[:, :, :z_t.shape[2]] + for blk in self.blocks: + h = blk(h, cond) + return self.out(h) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +class AttnRestorer(nn.Module): + """Deterministic restorer with global attention (Conformer-lite). Residual on the + degraded latent, like DetRestorer, so it inherits the identity-passthrough init.""" + def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, use_mix=True, + heads=4, t_max=64): + super().__init__() + self.use_mix = use_mix + self.stem_emb = nn.Embedding(n_stems, stem_emb) + in_ch = d * (2 if use_mix else 1) + self.inp = nn.Conv1d(in_ch, hidden, 1) + self.register_buffer("pos", _sinusoidal_pos(t_max, hidden), persistent=False) + self.blocks = nn.ModuleList([ConformerBlock(hidden, stem_emb, heads) for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) # start at identity (deg passthrough) + + def forward(self, deg, stem_id, mix=None): + cond = self.stem_emb(stem_id) # [B,stem_emb] + parts = [deg, mix] if (self.use_mix and mix is not None) else [deg] + h = self.inp(torch.cat(parts, dim=1)) # [B,hidden,T] + h = h + self.pos[:, :, :h.shape[2]] + for blk in self.blocks: + h = blk(h, cond) + return deg + self.out(h) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +class LatentDiscriminator(nn.Module): + """Spectral-norm conv critic over SAME latents [B,256,T] -> per-frame logits (PatchGAN-style) + + feature taps (for the stable feature-matching loss). Latent-domain = cheap, no decode + ('the latent is the mirror'). Judges full-mix realism in the remix-GAN fine-tune stage.""" + def __init__(self, d=256, hidden=256, depth=4, k=5): + super().__init__() + from torch.nn.utils import spectral_norm as SN + self.blocks = nn.ModuleList() + c = d + for _ in range(depth): + self.blocks.append(nn.Sequential(SN(nn.Conv1d(c, hidden, k, padding=k // 2)), + nn.LeakyReLU(0.2))) + c = hidden + self.head = SN(nn.Conv1d(c, 1, 1)) + + def forward(self, x, return_feats=False): + feats = []; h = x + for b in self.blocks: + h = b(h); feats.append(h) + logit = self.head(h) # [B,1,T] per-frame + return (logit, feats) if return_feats else logit + + def num_params(self): + return sum(p.numel() for p in self.parameters()) + + +def _sinusoidal_pos(t_max: int, dim: int) -> torch.Tensor: + pos = torch.arange(t_max).float()[:, None] + half = dim // 2 + freqs = torch.exp(-math.log(10000) * torch.arange(half).float() / max(half - 1, 1)) + pe = torch.zeros(t_max, dim) + pe[:, 0::2] = torch.sin(pos * freqs)[:, :pe[:, 0::2].shape[1]] + pe[:, 1::2] = torch.cos(pos * freqs)[:, :pe[:, 1::2].shape[1]] + return pe.t()[None] # [1,dim,t_max] + + +class DetRestorer(nn.Module): + """Deterministic restorer: predict clean stem latent from degraded stem (+ mix) + stem id. + Validated > generative for reference-matching; mix-conditioning helps drums most. + Predicts a residual on the degraded latent (small, aligned move).""" + def __init__(self, d=256, hidden=384, depth=5, n_stems=4, stem_emb=32, use_mix=True): + super().__init__() + self.use_mix = use_mix + self.stem_emb = nn.Embedding(n_stems, stem_emb) + in_ch = d * (2 if use_mix else 1) + stem_emb + self.inp = nn.Conv1d(in_ch, hidden, 1) + self.blocks = nn.ModuleList([ + nn.Sequential(nn.Conv1d(hidden, hidden, 5, padding=2), nn.GELU(), nn.Conv1d(hidden, hidden, 1)) + for _ in range(depth)]) + self.out = nn.Conv1d(hidden, d, 1) + nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) # start at identity (deg passthrough) + + def forward(self, deg, stem_id, mix=None): + em = self.stem_emb(stem_id)[:, :, None].expand(-1, -1, deg.shape[2]) + parts = [deg, mix, em] if (self.use_mix and mix is not None) else [deg, em] + h = self.inp(torch.cat(parts, dim=1)) + for blk in self.blocks: + h = h + blk(h) + return deg + self.out(h) + + def num_params(self): + return sum(p.numel() for p in self.parameters()) diff --git a/restoflow/next_moves_watch.py b/restoflow/next_moves_watch.py new file mode 100644 index 0000000000000000000000000000000000000000..e50299397f29c8cd2099c36f6d2d847a1cafac0b --- /dev/null +++ b/restoflow/next_moves_watch.py @@ -0,0 +1,143 @@ +"""Autonomous decision watcher for the two 'lighter moves' + capstone backup. + +Waits for: + - Move 1 (hi-adv GAN) FAD eval -> restoflow_runs/_logs/_hiadv_eval.log ("DONE") + - Move 2 (flowattn restorer) -> restoflow_runs/restorer_flowattn_v1/... (final eval in its log) +Parses the FADs, applies the decision tree, writes a verdict file, and โ€” ONLY if BOTH moves fail โ€” +auto-launches the capstone backup (remix_gan --train-gen: joint restorer+generator adversarial co-train, +no cache-regen needed, the safe single-command capstone). Promotion + HF push stay human-gated (recorded +as a RECOMMENDATION), so publishing never fires unattended on a parse. + +Baselines to beat (from quality_probe n=200): + REMIX latent-FAD: w003=5.561, w003_gan(w.05)=4.949 | drums: w003=3.221, w003_gan=3.343 +""" +from __future__ import annotations +import re, subprocess, time +from pathlib import Path + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") +LOGS = BASE / "restoflow_runs/_logs" +HIADV_EVAL = LOGS / "_hiadv_eval.log" +FLOW_LOG = LOGS / "restorer_flowattn_v1.log" +VERDICT = LOGS / "_next_moves_verdict.txt" +EE = "/home/ksoil/.conda/envs/ksoil_encoders/bin/python" + +W003_GAN_REMIX = 4.949 +W003_GAN_DRUMS = 3.343 +DRUMS_TOL = 0.10 # allow hi-adv drums to be no worse than w003_gan + tol + + +def wait_for(pred, timeout_s=10800, poll=60): + for _ in range(timeout_s // poll): + if pred(): + return True + time.sleep(poll) + return pred() + + +def parse_hiadv(): + """quality_probe table: 'REMIX ' -> last col = hiadv.""" + if not HIADV_EVAL.exists(): + return None + txt = HIADV_EVAL.read_text() + out = {} + for stem in ("REMIX", "drums"): + m = re.search(rf"^{stem}\s+([0-9.]+)\s+([0-9.]+)\s+([0-9.]+)\s+([0-9.]+)", txt, re.M) + if m: + out[stem] = float(m.group(4)) # hi-adv column (4th run-number) + return out or None + + +def parse_flow(): + """flowattn eval: per-stem 'drums | ->' and 'FAD-REMIX | ->'. Take the LAST block.""" + if not FLOW_LOG.exists(): + return None + txt = FLOW_LOG.read_text() + out = {} + for key, pat in (("drums", r"drums\s*\|\s*([0-9.]+)->([0-9.]+)"), + ("REMIX", r"FAD-REMIX\s*\|\s*([0-9.]+)->([0-9.]+)")): + ms = list(re.finditer(pat, txt)) + if ms: + out[key] = (float(ms[-1].group(1)), float(ms[-1].group(2))) # (in, out) + return out or None + + +def gpu_free(idx, tries=60): + for _ in range(tries): + try: + u = subprocess.check_output( + ["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader,nounits", "-i", str(idx)], + text=True).strip() + if int(u) < 1500: + return True + except Exception: + pass + time.sleep(30) + return False + + +def launch_capstone(): + """Joint co-train: restorer(w003_gan) + generator(gen_distvar_baseline) adversarial, generator UNFROZEN. + The single-command capstone (no cache regen). Conservative: same de-risked GAN, --train-gen + small w_adv.""" + if not gpu_free(0): + return "capstone NOT launched (cuda:0 never freed)" + log = LOGS / "remix_gan_cotrain.log" + cmd = [EE, "-m", "restoflow.remix_gan", + "--restorer", str(BASE / "restoflow_runs/restorer_attn_w003_gan"), + "--generator", str(BASE / "restoflow_runs/gen_distvar_baseline"), + "--out-dir", str(BASE / "restoflow_runs/remix_gan_cotrain"), + "--epochs", "15", "--batch", "24", "--adv-warmup", "400", + "--r1", "0.1", "--lr-d", "1e-4", "--lr-g", "1e-5", "--w-adv", "0.05", "--train-gen"] + env = {"CUDA_VISIBLE_DEVICES": "0", "PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True", + "PATH": "/usr/bin:/bin"} + import os + env = {**os.environ, **env} + subprocess.Popen(cmd, stdout=open(log, "w"), stderr=subprocess.STDOUT, env=env, cwd=str(BASE)) + return f"CAPSTONE LAUNCHED (remix_gan --train-gen) -> {log.name}" + + +def main(): + wait_for(lambda: HIADV_EVAL.exists() and "DONE" in HIADV_EVAL.read_text()) + # flowattn done: train.py prints 'done. best REMIX=...' as its LAST line, AFTER the ep50 eval+FAD. + # (Do NOT trigger on 'epoch 50/50' โ€” that prints before the final eval, so parse_flow would read the + # stale ep40 block.) This guarantees the ep50 FAD is flushed before we parse. + wait_for(lambda: FLOW_LOG.exists() and "done. best REMIX" in FLOW_LOG.read_text()) + time.sleep(15) # flush + + hi = parse_hiadv() or {} + fl = parse_flow() or {} + hi_remix = hi.get("REMIX"); hi_drums = hi.get("drums") + fl_drums = fl.get("drums"); fl_remix = fl.get("REMIX") + + hiadv_win = (hi_remix is not None and hi_remix < W003_GAN_REMIX + and (hi_drums is None or hi_drums <= W003_GAN_DRUMS + DRUMS_TOL)) + # flowattn judged INTERNALLY (its eval harness has a different FAD scale than quality_probe, so an + # absolute compare vs 3.343 is unfair): require out3.34 = x0.66); allow a little slack -> x0.72. + fl_ratio = (fl_drums[1] / fl_drums[0]) if fl_drums and fl_drums[0] else None + flow_drums_win = (fl_ratio is not None and fl_ratio < 0.72) + + lines = [f"=== NEXT-MOVES VERDICT {time.strftime('%F %T')} ===", + f"baselines: REMIX w003_gan={W003_GAN_REMIX} drums w003_gan={W003_GAN_DRUMS}", + f"MOVE1 hi-adv: REMIX-FAD={hi_remix} drums-FAD={hi_drums} -> {'WIN' if hiadv_win else 'no improvement'}", + f"MOVE2 flowattn: drums in->out={fl_drums} (ratio={fl_ratio}) REMIX in->out={fl_remix} -> " + f"{'WIN (drums, ratio<0.72)' if flow_drums_win else 'FAIL (drums reduction < deterministic)'}", ""] + + action = "" + if hiadv_win: + lines.append(f"RECOMMEND PROMOTE+PUSH: remix_gan_hiadv (REMIX {hi_remix} < {W003_GAN_REMIX}). " + "Rename -> restorer_attn_w003_gan_hi, set app+deploy default, re-push. [human-gated]") + if flow_drums_win: + lines.append(f"ADOPT FLOWATTN FOR DRUMS: drums-out {fl_drums[1]} < {W003_GAN_DRUMS}. " + "Wire hybrid (flowattn drums + deterministic/GAN others) in app inference. [human-gated]") + if not hiadv_win and not flow_drums_win: + lines.append("BOTH MOVES FAILED -> launching CAPSTONE backup (joint co-train).") + action = launch_capstone() + lines.append(action) + + VERDICT.write_text("\n".join(lines) + "\n") + print("\n".join(lines)) + + +if __name__ == "__main__": + main() diff --git a/restoflow/quality_probe.py b/restoflow/quality_probe.py new file mode 100644 index 0000000000000000000000000000000000000000..8a2dad5c3d928befb4869567762447e1df26e2c0 --- /dev/null +++ b/restoflow/quality_probe.py @@ -0,0 +1,119 @@ +"""Restorer QUALITY probe (orthogonal to REMIX's reference-matching floor): + (1) latent-FAD โ€” distributional distance to clean in SAME-latent space (cheap, no decode). + (2) spectral-dullness โ€” decoded HF-energy ratio + spectral centroid vs clean ("dull" = HF deficit + / lower centroid). Measured in DECODED space (latent energy-ratio misled us before). +Per stem + REMIX, for {degraded, } vs clean. CPU by default (won't disturb training GPUs). + +Run: python -m restoflow.quality_probe --runs restorer_attn_v1,restorer_attn_distvar_w003 --device cpu +""" +from __future__ import annotations +import argparse, random +from collections import defaultdict +import numpy as np, torch + +from .config import Cfg, STEMS, STEM_ID +from . import data as D, eval as E, fad +from .model import AttnRestorer, DetRestorer + +BASE = "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra" + + +def load_run(run, dev): + ck = torch.load(f"{BASE}/restoflow_runs/{run}/ckpt_best.pt", map_location="cpu"); c = ck["cfg"] + if c.get("model_kind") == "attn": + m = AttnRestorer(256, c["hidden"], c["depth"], n_stems=4, stem_emb=c["stem_emb_dim"], + use_mix=c["use_mix"], heads=c.get("heads", 4)) + else: + m = DetRestorer(256, c["hidden"], c["depth"], n_stems=4, stem_emb=c["stem_emb_dim"], use_mix=c["use_mix"]) + m.load_state_dict(ck["model"]); m.eval().to(dev) + return m, torch.load(f"{BASE}/restoflow_runs/{run}/norm_stats.pt", map_location="cpu") + + +def spectral(wav, sr, hf_hz=4000.0): + """wav [2,S] -> (HF-energy ratio, spectral centroid Hz) on the mono mix.""" + x = wav.mean(0) + S = torch.stft(x, 2048, hop_length=512, return_complex=True).abs() # [F,Tt] + freqs = torch.linspace(0, sr / 2, S.shape[0]) + e = S.sum(1) # energy per freq bin + tot = e.sum() + 1e-9 + hf = e[freqs >= hf_hz].sum() / tot + centroid = (freqs * e).sum() / tot + return hf.item(), centroid.item() + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--runs", default="restorer_attn_v1,restorer_attn_distvar_w003") + p.add_argument("--n-pairs", type=int, default=200) # latent-FAD (cheap) + p.add_argument("--n-decode", type=int, default=48) # spectral (decode, CPU) + p.add_argument("--device", default="cpu") + a = p.parse_args() + dev = a.device; random.seed(0) + runs = [r for r in a.runs.split(",") if r] + conds = ["degraded"] + runs # vs clean + cfg = Cfg(val_cache_roots=(f"{BASE}/demucs_results_val_full",), use_mix=True, model_kind="attn", T=32, device=dev) + + models = {r: load_run(r, dev) for r in runs} + items = D._index(cfg, cfg.val_cache_roots, "val") + by_pair = defaultdict(list) + for it in items: by_pair[it["pair_id"]].append(it) + pairs = list(by_pair.items()); random.shuffle(pairs) + freq = defaultdict(int) + for _, its in pairs: + for it in its: freq[it["stem"]] += 1 + pairs.sort(key=lambda kv: min(freq[it["stem"]] for it in kv[1])) # rare-stem first + pairs = pairs[:a.n_pairs] + + lat = {c: defaultdict(list) for c in ["clean"] + conds} # cond->stem->[ [256,T] ] + rlat = {c: [] for c in ["clean"] + conds} # remix per pair + with torch.inference_mode(): + for (pid, its) in pairs: + psum = {c: None for c in ["clean"] + conds} + for it in its: + s = it["stem"]; sid = torch.tensor([STEM_ID[s]], device=dev) + raw_c = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) + raw_d = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) + raw_m = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) + vals = {"clean": raw_c, "degraded": raw_d} + for r in runs: + m, st = models[r] + sd = ((raw_d - st[s]["mu"][:, None]) / st[s]["sd"][:, None]).to(dev)[None] + sm = ((raw_m - st["__mix__"]["mu"][:, None]) / st["__mix__"]["sd"][:, None]).to(dev)[None] + out = m(sd, sid, sm)[0] + vals[r] = E._destd(out, st, s).cpu() + for c in ["clean"] + conds: + lat[c][s].append(vals[c]) + psum[c] = vals[c] if psum[c] is None else psum[c] + vals[c] + for c in ["clean"] + conds: + if psum[c] is not None: rlat[c].append(psum[c]) + + # ---------- (1) latent-FAD ---------- + print(f"\n=== latent-FAD vs clean (n_pairs={len(pairs)}) โ€” distributional, lower=closer ===") + print(f"{'stem':7} " + " ".join(f"{c[:14]:>14}" for c in conds)) + for s in STEMS: + if not lat["clean"][s]: continue + cl = fad.latent_set(lat["clean"][s]) + row = [fad.fad(cl, fad.latent_set(lat[c][s])) for c in conds] + print(f"{s:7} " + " ".join(f"{v:14.3f}" for v in row)) + clr = fad.latent_set(rlat["clean"]) + print(f"{'REMIX':7} " + " ".join(f"{fad.fad(clr, fad.latent_set(rlat[c])):14.3f}" for c in conds)) + + # ---------- (2) spectral-dullness (decode subset) ---------- + sa, sr = E.load_same(cfg) + def dec(latents): + return [E._decode(sa, l.to(dev)) for l in latents] + print(f"\n=== spectral-dullness (decoded, n={a.n_decode}) โ€” HF-ratio & centroid; dull = below clean ===") + print(f"{'stem':7} {'metric':9} " + " ".join(f"{c[:12]:>12}" for c in ['clean'] + conds)) + for s in STEMS: + if not lat["clean"][s]: continue + n = min(a.n_decode, len(lat["clean"][s])) + hf = {}; ce = {} + for c in ['clean'] + conds: + feats = [spectral(w, sr) for w in dec(lat[c][s][:n])] + hf[c] = float(np.mean([f[0] for f in feats])); ce[c] = float(np.mean([f[1] for f in feats])) + print(f"{s:7} {'HF-ratio':9} " + " ".join(f"{hf[c]:12.4f}" for c in ['clean'] + conds)) + print(f"{'':7} {'centroid':9} " + " ".join(f"{ce[c]:12.0f}" for c in ['clean'] + conds)) + + +if __name__ == "__main__": + main() diff --git a/restoflow/remix_gan.py b/restoflow/remix_gan.py new file mode 100644 index 0000000000000000000000000000000000000000..decdcb34440339d60cecbed859a89117da3ef5d7 --- /dev/null +++ b/restoflow/remix_gan.py @@ -0,0 +1,283 @@ +"""FINAL stage: remix-coupled, recon+distribution-ANCHORED, LIGHT-adversarial fine-tune. + +Restorer fixes restoration-class stems, generator fills generation-class stems, sum -> reconstructed +mix latent. A latent discriminator judges real clean-mix vs reconstructed-mix. BOTH models are the +"generators" co-trained to make a coherent realistic full mix. De-risked by construction: + - recon (restorer MSE) + var-match distribution anchor DOMINATE; adversarial is a small ramped aux + with a kill-switch -> worst case == the non-adversarial model, never "diverged garbage". + - feature-matching (stable) + hinge + R1 grad penalty + warm-started D + grad-accum for OOM. +Latent-domain (no decode) = cheap, CPU-smokeable. NOT launched automatically. + +Smoke: python -m restoflow.remix_gan --smoke --device cpu +""" +from __future__ import annotations +import argparse, csv, random, time +from collections import defaultdict +from pathlib import Path +import torch, torch.nn as nn, torch.nn.functional as F +from torch.utils.data import Dataset, DataLoader + +from .config import Cfg, STEMS, STEM_ID +from . import data as D, router, distloss as DL +from .model import AttnRestorer, DetRestorer, CondFlow, AttnCondFlow, LatentDiscriminator + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") + + +# ---------------- pair-grouped data ---------------- +class PairData(Dataset): + """Per pair: clean[4], deg[4] (raw latents, 0 if absent), mix, present[4], role[4] + (0 none / 1 restoration / 2 generation).""" + def __init__(self, roots, cfg, max_pairs=None): + self.cfg = cfg; self.T = cfg.T + rows = defaultdict(dict) + for root in roots: + meta = Path(root) / "metadata.csv" + if not meta.exists(): + continue + for r in csv.DictReader(open(meta)): + if r.get("stem") not in STEMS: + continue + rows[(root, r["pair_id"])][r["stem"]] = r + self.items = [] + for (root, pid), stemrows in rows.items(): + mix = Path(root) / pid / "latents" / "degraded_mix.pt" + if not mix.exists(): + continue + self.items.append((root, pid, stemrows, str(mix))) + random.Random(0).shuffle(self.items) + if max_pairs: + self.items = self.items[:max_pairs] + + def __len__(self): return len(self.items) + + def __getitem__(self, i): + root, pid, stemrows, mixp = self.items[i] + T = self.T + clean = torch.zeros(len(STEMS), 256, T); deg = torch.zeros(len(STEMS), 256, T) + present = torch.zeros(len(STEMS), dtype=torch.bool); role = torch.zeros(len(STEMS), dtype=torch.long) + for s, r in stemrows.items(): + si = STEM_ID[s] + cl = r.get("clean_latent_path") or "" + if cl: + p = Path(root) / cl + if p.exists(): + clean[si] = D._fit_T(torch.load(p, map_location="cpu").float(), T); present[si] = True + if not present[si]: + continue + if router.is_restoration(r, self.cfg.rms_thr, self.cfg.peak_thr, self.cfg.max_drop_db): + dp = r.get("deg_latent_path") or "" + if dp and (Path(root) / dp).exists(): + deg[si] = D._fit_T(torch.load(Path(root) / dp, map_location="cpu").float(), T) + role[si] = 1 # restoration + else: + role[si] = 2 # generation (clean present, deg gutted) + mix = D._fit_T(torch.load(mixp, map_location="cpu").float(), T) + return {"clean": clean, "deg": deg, "mix": mix, "present": present, "role": role} + + +# ---------------- helpers ---------------- +def _std(x, mu, sd): return (x - mu[:, None]) / sd[:, None] +def _destd(x, mu, sd): return x * sd[:, None] + mu[:, None] + + +def load_restorer(run, dev): + ck = torch.load(run / "ckpt_best.pt", map_location="cpu"); c = ck["cfg"] + cls = AttnRestorer if c.get("model_kind") == "attn" else DetRestorer + kw = dict(stem_emb=c["stem_emb_dim"], use_mix=c["use_mix"]) + if c.get("model_kind") == "attn": kw["heads"] = c.get("heads", 4) + m = cls(256, c["hidden"], c["depth"], n_stems=len(STEMS), **kw).to(dev) + m.load_state_dict(ck["model"]) + return m, torch.load(run / "norm_stats.pt", map_location="cpu"), c + + +def load_generator(run, dev): + ck = torch.load(run / "ckpt_best.pt", map_location="cpu"); a = ck["args"] + cls = AttnCondFlow if a.get("arch") == "attn" else CondFlow + kw = dict(stem_emb=Cfg().stem_emb_dim) + if a.get("arch") == "attn": kw["heads"] = a.get("heads", 8) + m = cls(256, a["hidden"], a["depth"], n_stems=len(STEMS), **kw).to(dev) + m.load_state_dict(ck["model"]) + return m, ck["stats"], a + + +def sample_grad(gen, ctx, sid, steps): + """Few-step Euler flow sampler WITH gradients (training-time; CFG off).""" + z = torch.randn_like(ctx); dt = 1.0 / steps + for i in range(steps): + t = torch.full((ctx.shape[0],), i * dt, device=ctx.device) + z = z + dt * gen(z, t, ctx, sid) + return z + + +def r1_penalty(D, real): + """Gradient penalty, MEAN over latent dims (not sum). Sum-over-8192-dims made this ~O(80) and, + with the old r1=1.0 every step, it dwarfed the hinge (~2) -> D collapsed to a constant (Dacc 0.50). + Mean-normalization makes it scale-invariant (~O(grad_per_elem^2)); spectral-norm already bounds D's + Lipschitz so this is a light touch, applied lazily.""" + real = real.detach().requires_grad_(True) + out = D(real).sum() + g, = torch.autograd.grad(out, real, create_graph=True) + return g.pow(2).flatten(1).mean(1).mean() + + +# ---------------- train ---------------- +def main(): + p = argparse.ArgumentParser() + p.add_argument("--restorer", default=str(BASE / "restoflow_runs/restorer_attn_v1")) + p.add_argument("--generator", default=str(BASE / "restoflow_runs_archive/gen_v2")) + p.add_argument("--out-dir", default=str(BASE / "restoflow_runs/remix_gan_v1")) + p.add_argument("--cache-roots", default=str(BASE / "demucs_results_all")) + p.add_argument("--device", default="cuda") + p.add_argument("--epochs", type=int, default=20) + p.add_argument("--batch", type=int, default=16) + p.add_argument("--accum", type=int, default=4, help="grad-accum micro-batches (OOM -> small batch, big effective)") + p.add_argument("--lr-g", type=float, default=2e-5); p.add_argument("--lr-d", type=float, default=1e-4) # TTUR + p.add_argument("--gen-steps", type=int, default=2) + p.add_argument("--w-recon", type=float, default=1.0) + p.add_argument("--w-dist", type=float, default=0.02) + p.add_argument("--w-fm", type=float, default=2.0) # feature-matching = the stable workhorse + p.add_argument("--w-adv", type=float, default=0.05) # SMALL adversarial (ramped) + p.add_argument("--adv-warmup", type=int, default=200) # D-only warmup steps before adv on G + p.add_argument("--train-gen", action="store_true", + help="co-train the generator too (risky: adv-through-sampler). Default: generator " + "FROZEN (sampled no-grad), only restorer+D adapt first โ€” the safe stage.") + p.add_argument("--r1", type=float, default=0.1) # mean-normalized R1 (spectral-norm already bounds D) + p.add_argument("--d-reg-every", type=int, default=16) # lazy R1: every N steps (cost; magnitude is fixed by mean-norm) + p.add_argument("--num-workers", type=int, default=4) + p.add_argument("--smoke", action="store_true") + a = p.parse_args() + if a.smoke: + a.epochs, a.batch, a.accum, a.num_workers, a.adv_warmup = 1, 3, 1, 0, 2 + dev = a.device + out = Path(a.out_dir); out.mkdir(parents=True, exist_ok=True) + cfg = Cfg(device=dev, T=32) + + rest, rstats, rcfg = load_restorer(Path(a.restorer), dev) + gen, gstats, gargs = load_generator(Path(a.generator), dev) + disc = LatentDiscriminator().to(dev) + print(f"[gan] restorer {sum(p.numel() for p in rest.parameters())/1e6:.1f}M " + f"generator {sum(p.numel() for p in gen.parameters())/1e6:.1f}M disc {disc.num_params()/1e6:.1f}M") + + rmu = {s: rstats[s]["mu"].to(dev) for s in STEMS}; rsd = {s: rstats[s]["sd"].to(dev) for s in STEMS} + mmu, msd = rstats["__mix__"]["mu"].to(dev), rstats["__mix__"]["sd"].to(dev) + gtm = {s: gstats["tgt"][s][0].to(dev) for s in STEMS}; gts = {s: gstats["tgt"][s][1].to(dev) for s in STEMS} + gcm = {s: gstats["ctx"][s][0].to(dev) for s in STEMS}; gcs = {s: gstats["ctx"][s][1].to(dev) for s in STEMS} + + ds = PairData([x for x in a.cache_roots.split(",") if x], cfg, max_pairs=60 if a.smoke else None) + loader = DataLoader(ds, batch_size=a.batch, shuffle=True, drop_last=True, num_workers=a.num_workers) + print(f"[gan] {len(ds)} pairs") + gparams = list(rest.parameters()) + (list(gen.parameters()) if a.train_gen else []) + if not a.train_gen: + gen.eval() + for p_ in gen.parameters(): p_.requires_grad_(False) + optG = torch.optim.AdamW(gparams, a.lr_g, betas=(0.5, 0.9)) + optD = torch.optim.AdamW(disc.parameters(), a.lr_d, betas=(0.5, 0.9)) + + def build_remix(b): + """-> (recon_mix[B,256,T], real_mix[B,256,T], recon_loss). Restorer for role1, generator role2.""" + clean = b["clean"].to(dev); deg = b["deg"].to(dev); mix = b["mix"].to(dev) + present = b["present"].to(dev); role = b["role"].to(dev) + B = clean.shape[0] + std_mix = _std(mix, mmu, msd) + restored = torch.zeros_like(clean); recon_l = clean.new_zeros(()) + nrec = 0 + for si, s in enumerate(STEMS): + m = role[:, si] == 1 + if m.any(): + sid = torch.full((int(m.sum()),), si, device=dev) + sd_in = _std(deg[m, si], rmu[s], rsd[s]) + out = rest(sd_in, sid, std_mix[m]) + restored[m, si] = _destd(out, rmu[s], rsd[s]) + recon_l = recon_l + F.mse_loss(out, _std(clean[m, si], rmu[s], rsd[s])); nrec += 1 + recon_l = recon_l / max(nrec, 1) + # generation stems: context = sum of restored present-others, normalize by gen ctx stats. + # ALSO compute the generator's own FM loss toward the clean target = ANCHOR (keeps the + # generator a valid flow under adversarial pressure; prevents the drift/explosion). + gen_out = torch.zeros_like(clean); gen_fm = clean.new_zeros(()); ngen = 0 + for si, s in enumerate(STEMS): + m = role[:, si] == 2 + if m.any(): + others = present.clone(); others[:, si] = False + ctx_raw = (restored * others[:, :, None, None].float())[m].sum(1) # [n,256,T] sum others + sid = torch.full((int(m.sum()),), si, device=dev) + ctx = _std(ctx_raw, gcm[s], gcs[s]) + if a.train_gen: # co-train: FM anchor + grad sampler + tgt = _std(clean[m, si], gtm[s], gts[s]) + x0 = torch.randn_like(tgt); t = torch.rand(tgt.shape[0], device=dev) + z = (1 - t[:, None, None]) * x0 + t[:, None, None] * tgt + gen_fm = gen_fm + F.mse_loss(gen(z, t, ctx, sid), tgt - x0); ngen += 1 + g = sample_grad(gen, ctx, sid, a.gen_steps) + else: # frozen generator (safe stage) + with torch.no_grad(): + g = sample_grad(gen, ctx, sid, a.gen_steps) + gen_out[m, si] = _destd(g, gtm[s], gts[s]) + gen_fm = gen_fm / max(ngen, 1) + use = torch.where((role == 2)[:, :, None, None], gen_out, + torch.where((role == 1)[:, :, None, None], restored, clean)) + use = use * present[:, :, None, None].float() + recon_mix = use.sum(1) + real_mix = (clean * present[:, :, None, None].float()).sum(1) + return recon_mix, real_mix, recon_l, gen_fm + + step = 0 + for ep in range(a.epochs): + for b in loader: + recon_mix, real_mix, recon_l, gen_fm = build_remix(b) + real_in = _std(real_mix, mmu, msd); recon_in = _std(recon_mix, mmu, msd) # standardize D inputs + # ---- D step ---- + optD.zero_grad() + d_real = disc(real_in); d_fake = disc(recon_in.detach()) + lossD = F.relu(1 - d_real).mean() + F.relu(1 + d_fake).mean() + if a.r1 > 0 and step % a.d_reg_every == 0: + lossD = lossD + a.r1 * r1_penalty(disc, real_in) + if torch.isfinite(lossD): + lossD.backward() + torch.nn.utils.clip_grad_norm_(disc.parameters(), 1.0) + optD.step() + # ---- G step ---- + optG.zero_grad() + d_fake_logit, f_fake = disc(_std(recon_mix, mmu, msd), return_feats=True) + with torch.no_grad(): + _, f_real = disc(real_in, return_feats=True) + fm = sum(F.l1_loss(ff, fr) for ff, fr in zip(f_fake, f_real)) / len(f_fake) + dist = DL.mmd_loss(recon_mix, real_mix) # bounded -> stable mix anchor + adv = -d_fake_logit.mean() + dacc = ((d_real > 0).float().mean() + (d_fake < 0).float().mean()).item() / 2 + # KILL-SWITCH: adversarial only after warmup AND while D isn't dominating (Dacc<0.8). + # When D wins, adv/fm gradients blow up -> drop them; recon+gen_fm+mmd anchors carry on. + adv_on = step >= a.adv_warmup and dacc < 0.8 + w_adv = a.w_adv if adv_on else 0.0 + w_fm = a.w_fm if adv_on else 0.0 + lossG = a.w_recon * (recon_l + gen_fm) + a.w_dist * dist + w_fm * fm + w_adv * adv + if torch.isfinite(lossG): + lossG.backward() + torch.nn.utils.clip_grad_norm_(list(rest.parameters()) + list(gen.parameters()), 1.0) + optG.step() + else: + print(f" [skip] non-finite G loss @ step {step}") + step += 1 + if step % (1 if a.smoke else 50) == 0: + print(f"ep{ep} step{step} recon={recon_l.item():.4f} genfm={gen_fm.item():.4f} " + f"mmd={dist.item():.4f} fm={fm.item():.4f} adv={adv.item():.3f} | " + f"D={lossD.item():.3f} Dacc={dacc:.2f} dR={d_real.mean().item():+.2f} " + f"dF={d_fake.mean().item():+.2f} w_adv={w_adv}") + # save after each epoch in the STANDARD restorer run-dir schema (model+cfg+norm_stats), so the + # adversarially-refined restorer is a drop-in for quality_probe / select_best / app.py โ€” plus + # the discriminator for resuming. Recoverable + directly FAD-comparable vs the source restorer. + torch.save({"model": rest.state_dict(), "cfg": rcfg, "epoch": ep + 1, + "disc": disc.state_dict(), "gan_args": vars(a), + "source_restorer": Path(a.restorer).name}, out / "ckpt_best.pt") + torch.save(rstats, out / "norm_stats.pt") + if a.train_gen: # capture the CO-TRAINED generator as a standalone + gout = out.parent / f"{out.name}_gen" # generator run-dir (loadable by gen eval / app) + gout.mkdir(parents=True, exist_ok=True) + torch.save({"model": gen.state_dict(), "args": gargs, "stats": gstats, + "epoch": ep + 1, "co_trained_with": Path(a.restorer).name}, gout / "ckpt_best.pt") + print(f"done. refined restorer -> {out/'ckpt_best.pt'} (drop-in run-dir; FAD-compare vs " + f"{Path(a.restorer).name} via quality_probe --runs). " + f"kill-switch = drop w_adv if Dacc->1 (D wins) or quality/FAD regress.") + + +if __name__ == "__main__": + main() diff --git a/restoflow/router.py b/restoflow/router.py new file mode 100644 index 0000000000000000000000000000000000000000..48da527b51de4bd507e544558c08866d19058cca --- /dev/null +++ b/restoflow/router.py @@ -0,0 +1,44 @@ +"""Presence router: decide which (pair, stem) is a RESTORATION case (in scope). + +Restoration-only: we keep a stem pair iff the stem is present (rms OR peak) in BOTH +clean and degraded. generation / both_silent / deg_only are out of scope and skipped. +Works off metadata.csv columns, recomputing presence with hardened thresholds so we +do not need to re-cache when tuning the floor. +""" +from __future__ import annotations +import hashlib + + +def _f(row, key, default=None): + v = row.get(key, "") + try: + return float(v) + except (TypeError, ValueError): + return default + + +def present(rms_db, peak_db, rms_thr, peak_thr) -> bool: + if rms_db is None: + return False + if rms_db > rms_thr: + return True + return peak_db is not None and peak_db > peak_thr + + +def is_restoration(row, rms_thr=-40.0, peak_thr=-25.0, max_drop=15.0) -> bool: + """True iff clean AND degraded stem are present (hardened) AND the degraded stem is not + gutted (clean_rms - deg_rms <= max_drop). A gutted stem is generation, not restoration.""" + c = present(_f(row, "clean_rms_db"), _f(row, "clean_peak_db"), rms_thr, peak_thr) + d = present(_f(row, "deg_rms_db"), _f(row, "deg_peak_db"), rms_thr, peak_thr) + if not (c and d): + return False + cr, dr = _f(row, "clean_rms_db"), _f(row, "deg_rms_db") + if cr is not None and dr is not None and (cr - dr) > max_drop: + return False + return True + + +def pair_split(pair_id: str, val_frac: float, seed: int) -> str: + """Deterministic, stable train/val assignment keyed on pair_id (no leakage).""" + h = int(hashlib.md5(f"{seed}:{pair_id}".encode()).hexdigest(), 16) % 10_000 + return "val" if h < int(val_frac * 10_000) else "train" diff --git a/restoflow/select_best.py b/restoflow/select_best.py new file mode 100644 index 0000000000000000000000000000000000000000..371c4c287eb9497c8aee3d3c6474831bd9a516c4 --- /dev/null +++ b/restoflow/select_best.py @@ -0,0 +1,84 @@ +"""Auto-select the best restorer + best generator by latent-FAD (REMIX), and write _best.txt +(BEST_RESTORER / BEST_GEN) for the experimental GAN to consume. Decode-free. + +Restorers: ranked here by REMIX latent-FAD (distributional quality) with REMIX-STFT shown too. +Generators: read from the phase-2 gen-FAD comparison log (_gen_fad_compare.log). +""" +from __future__ import annotations +import argparse, glob, random, re +from collections import defaultdict +from pathlib import Path +import torch +from .config import Cfg, STEMS, STEM_ID +from . import data as D, fad +from .quality_probe import load_run + +BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") + + +def restorer_remix_fad(run, dev, pairs, cfg): + m, st = load_run(run, dev) + cl_set, rr_set = [], [] + with torch.inference_mode(): + for (_pid, its) in pairs: + clat = rlat = None + for it in its: + s = it["stem"]; sid = torch.tensor([STEM_ID[s]], device=dev) + rc = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) + rd = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) + rm = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) + sd = ((rd - st[s]["mu"][:, None]) / st[s]["sd"][:, None]).to(dev)[None] + sm = ((rm - st["__mix__"]["mu"][:, None]) / st["__mix__"]["sd"][:, None]).to(dev)[None] + rr = D._fit_T(_destd_cpu(m(sd, sid, sm)[0], st, s), cfg.T) + clat = rc if clat is None else clat + rc + rlat = rr if rlat is None else rlat + rr + if clat is not None: + cl_set.append(clat); rr_set.append(rlat) + return fad.fad(fad.latent_set(cl_set), fad.latent_set(rr_set)) + + +def _destd_cpu(z, stats, stem): + return (z * stats[stem]["sd"][:, None].to(z.device) + stats[stem]["mu"][:, None].to(z.device)).detach().cpu() + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--device", default="cuda"); p.add_argument("--n-pairs", type=int, default=150) + a = p.parse_args(); dev = a.device; random.seed(0) + cfg = Cfg(val_cache_roots=(str(BASE / "demucs_results_val_full"),), use_mix=True, model_kind="attn", T=32, device=dev) + items = D._index(cfg, cfg.val_cache_roots, "val") + by_pair = defaultdict(list) + for it in items: by_pair[it["pair_id"]].append(it) + pairs = list(by_pair.items()); random.shuffle(pairs); pairs = pairs[:a.n_pairs] + + rest_runs = [Path(d).name for d in sorted(glob.glob(str(BASE / "restoflow_runs/restorer_*"))) + if (Path(d) / "ckpt_best.pt").exists()] + print(f"=== restorer REMIX latent-FAD (lower=better) over {len(pairs)} pairs ===") + scores = {} + for r in rest_runs: + try: + scores[r] = restorer_remix_fad(r, dev, pairs, cfg); print(f" {r:32} {scores[r]:.3f}") + except Exception as e: + print(f" {r:32} FAILED {e}") + best_rest = min(scores, key=scores.get) if scores else "restorer_attn_v1" + + # generators: parse phase-2 comparison + gfad = {}; clog = BASE / "restoflow_runs/_gen_fad_compare.log" + if clog.exists(): + cur = None + for ln in clog.read_text().splitlines(): + mrun = re.search(r"#####\s+(\S+)", ln) + if mrun: cur = mrun.group(1) + mf = re.search(r"FAD-REMIX\s*\|\s*([0-9.]+)", ln) + if mf and cur: gfad[cur] = float(mf.group(1)) + print(f"\n=== generator REMIX-FAD (from _gen_fad_compare.log) ===") + for g, v in gfad.items(): print(f" {g:32} {v:.3f}") + best_gen = min(gfad, key=gfad.get) if gfad else "gen_distvar_baseline" + + out = BASE / "restoflow_runs/_best.txt" + out.write_text(f"BEST_RESTORER={best_rest}\nBEST_GEN={best_gen}\n") + print(f"\nBEST_RESTORER={best_rest}\nBEST_GEN={best_gen}\n-> {out}") + + +if __name__ == "__main__": + main() diff --git a/restoflow/train.py b/restoflow/train.py new file mode 100644 index 0000000000000000000000000000000000000000..b04156d16bcf7bf77e4e9a434806e64850cdb123 --- /dev/null +++ b/restoflow/train.py @@ -0,0 +1,151 @@ +"""Train the conditional flow restorer. Runnable: python -m restoflow.train --epochs 200 + +Stage A (separate+encode) is done by demucs.py; this consumes the latent cache. +Builds norm stats + train/val split first, then flow-matching training with periodic +per-stem + remix decoded eval and demo audio. +""" +from __future__ import annotations +import argparse, json, time +from dataclasses import replace +from pathlib import Path +import torch +from torch.utils.data import DataLoader + +from .config import Cfg +from . import data as D +from . import eval as E +from . import flow as F +from .model import CondFlow, DetRestorer, AttnRestorer, MixAttnCondFlow +from . import distloss as DL + + +def parse() -> Cfg: + c = Cfg(); p = argparse.ArgumentParser() + p.add_argument("--cache-roots", default=",".join(c.cache_roots)) + p.add_argument("--val-cache-roots", default=",".join(c.val_cache_roots)) + p.add_argument("--out-dir", default=c.out_dir) + p.add_argument("--epochs", type=int, default=c.epochs) + p.add_argument("--batch", type=int, default=c.batch) + p.add_argument("--lr", type=float, default=c.lr) + p.add_argument("--eval-every", type=int, default=c.eval_every) + p.add_argument("--device", default=c.device) + p.add_argument("--num-workers", type=int, default=c.num_workers) + p.add_argument("--sigma", type=float, default=c.sigma) + p.add_argument("--hidden", type=int, default=c.hidden) + p.add_argument("--depth", type=int, default=c.depth) + p.add_argument("--model-kind", default=c.model_kind, choices=["det", "attn", "flow", "flowattn"]) + p.add_argument("--heads", type=int, default=c.heads) + p.add_argument("--rms-thr", type=float, default=c.rms_thr) + p.add_argument("--peak-thr", type=float, default=c.peak_thr) + p.add_argument("--max-drop-db", type=float, default=c.max_drop_db) + p.add_argument("--eval-max-pairs", type=int, default=c.eval_max_pairs) + p.add_argument("--stem-weight-alpha", type=float, default=c.stem_weight_alpha) + p.add_argument("--energy-w", type=float, default=c.energy_w) + p.add_argument("--dist-loss", default=c.dist_loss, choices=["none", "var", "mmd", "moment"], + help="cheap latent distributional aux loss (fights deterministic dullness)") + p.add_argument("--dist-loss-w", type=float, default=c.dist_loss_w) + p.add_argument("--smoke", action="store_true", help="tiny run to verify the pipeline") + a = p.parse_args() + cr = tuple(x for x in a.cache_roots.split(",") if x) + vr = tuple(x for x in a.val_cache_roots.split(",") if x) + cfg = replace(c, cache_roots=cr, val_cache_roots=vr, out_dir=a.out_dir, epochs=a.epochs, + batch=a.batch, lr=a.lr, eval_every=a.eval_every, device=a.device, + num_workers=a.num_workers, sigma=a.sigma, hidden=a.hidden, depth=a.depth, + rms_thr=a.rms_thr, peak_thr=a.peak_thr, max_drop_db=a.max_drop_db, + eval_max_pairs=a.eval_max_pairs, stem_weight_alpha=a.stem_weight_alpha, + energy_w=a.energy_w, model_kind=a.model_kind, heads=a.heads, + dist_loss=a.dist_loss, dist_loss_w=a.dist_loss_w) + if a.smoke: + cfg = replace(cfg, epochs=2, eval_every=1, demo_pairs=2, num_workers=0) + return cfg + + +def stem_weights(cfg, items, device): + """Per-stem loss weights, mean-normalized to 1. alpha=0 -> uniform (v3); + alpha=1 -> full inverse-frequency (bass/drums no longer drowned by other/vocals).""" + from collections import Counter + cnt = Counter(i["stem"] for i in items) + freq = torch.tensor([max(1, cnt[s]) for s in cfg.stems], dtype=torch.float) + inv = (freq.sum() / (len(cfg.stems) * freq)) ** cfg.stem_weight_alpha + w = inv / inv.mean() + print(f"[loss] stem weights (alpha={cfg.stem_weight_alpha}): " + f"{dict(zip(cfg.stems, [round(x, 3) for x in w.tolist()]))} energy_w={cfg.energy_w}") + return w.to(device) + + +def det_loss(pred, clean, sid, stem_w, energy_w): + """Stem-weighted MSE + relative-energy(norm)-matching term, all in standardized space.""" + se = ((pred - clean) ** 2).flatten(1).mean(1) # [B] per-sample MSE + pn = pred.flatten(1).norm(dim=1); cn = clean.flatten(1).norm(dim=1) + en = ((pn - cn) / (cn + 1e-6)) ** 2 # [B] relative energy mismatch + w = stem_w[sid] # [B] + return (w * (se + energy_w * en)).sum() / (w.sum() + 1e-8) + + +def main(): + cfg = parse() + torch.manual_seed(0) + Path(cfg.out_dir).mkdir(parents=True, exist_ok=True) + json.dump(cfg.__dict__, open(Path(cfg.out_dir) / "config.json", "w"), indent=2, default=str) + + stats = D.load_or_build_stats(cfg) + train_ds = D.LatentRestore(cfg, "train", stats) + val_ds = D.LatentRestore(cfg, "val", stats) + from collections import Counter + print(f"[data] train={len(train_ds)} {dict(Counter(i['stem'] for i in train_ds.items))}") + print(f"[data] val ={len(val_ds)} {dict(Counter(i['stem'] for i in val_ds.items))}") + if len(train_ds) == 0: + raise SystemExit("No restoration items found โ€” check cache_roots / metadata.csv / thresholds.") + + loader = DataLoader(train_ds, batch_size=cfg.batch, shuffle=True, drop_last=True, + num_workers=cfg.num_workers, pin_memory=True) + if cfg.model_kind == "det": + model = DetRestorer(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim, use_mix=cfg.use_mix).to(cfg.device) + elif cfg.model_kind == "attn": + model = AttnRestorer(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim, use_mix=cfg.use_mix, heads=cfg.heads).to(cfg.device) + elif cfg.model_kind == "flowattn": # stochastic attention restorer + mix conditioning + model = MixAttnCondFlow(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim, heads=cfg.heads, use_mix=cfg.use_mix).to(cfg.device) + else: + model = CondFlow(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), + stem_emb=cfg.stem_emb_dim).to(cfg.device) + print(f"[model] {cfg.model_kind} params={model.num_params()/1e6:.2f}M use_mix={cfg.use_mix}") + stem_w = stem_weights(cfg, train_ds.items, cfg.device) + opt = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.wd) + sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, cfg.epochs * max(1, len(loader))) + sa, sr = E.load_same(cfg) + + best_remix = float("inf") + for epoch in range(cfg.epochs): + model.train(); t0 = time.time(); tot = 0.0 + for b in loader: + deg = b["deg"].to(cfg.device); clean = b["clean"].to(cfg.device); sid = b["stem_id"].to(cfg.device) + mix = b["mix"].to(cfg.device) if cfg.use_mix else None + if cfg.model_kind in ("det", "attn"): + pred = model(deg, sid, mix) + loss = det_loss(pred, clean, sid, stem_w, cfg.energy_w) + if cfg.dist_loss_w > 0: # latent distributional term -> de-dull + loss = loss + cfg.dist_loss_w * DL.dist_term(cfg.dist_loss, pred, clean, sid) + elif cfg.model_kind == "flowattn": + loss = F.fm_loss_mix(model, deg, clean, mix, sid, cfg.sigma) + else: + loss = F.fm_loss(model, deg, clean, sid, cfg.sigma) + opt.zero_grad(); loss.backward(); opt.step(); sched.step() + tot += loss.item() + print(f"epoch {epoch+1}/{cfg.epochs} loss={tot/max(1,len(loader)):.4f} ({time.time()-t0:.0f}s)") + if (epoch + 1) % cfg.eval_every == 0 or epoch == cfg.epochs - 1: + metrics = E.evaluate(cfg, model, stats, sa, sr, epoch=epoch + 1) + ckpt = {"model": model.state_dict(), "cfg": cfg.__dict__, "epoch": epoch + 1, "metrics": metrics} + torch.save(ckpt, Path(cfg.out_dir) / "ckpt.pt") # latest + remix = float(metrics.get("remix", {}).get("stft_out", float("inf"))) + if remix < best_remix: # keep the best (det model overfits late) + best_remix = remix + torch.save(ckpt, Path(cfg.out_dir) / "ckpt_best.pt") + print(f" ** new best REMIX={remix:.3f} @ epoch {epoch+1} -> ckpt_best.pt") + print(f"done. best REMIX={best_remix:.3f}") + + +if __name__ == "__main__": + main() diff --git a/restoflow_runs/gen_advramp_v1/ckpt_best.pt b/restoflow_runs/gen_advramp_v1/ckpt_best.pt new file mode 100644 index 0000000000000000000000000000000000000000..e63bf5e7693dc447c2f0f8874279b6699e24084e --- /dev/null +++ b/restoflow_runs/gen_advramp_v1/ckpt_best.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3fd1d935131a62500801bcf934f5075275b6be16fa5374e30e3b41e3d12c1de +size 41804850 diff --git a/restoflow_runs/gen_advramp_v1/config.json b/restoflow_runs/gen_advramp_v1/config.json new file mode 100644 index 0000000000000000000000000000000000000000..51f4addcb66848da6f90e9fc700d6bc9e5324fc7 --- /dev/null +++ b/restoflow_runs/gen_advramp_v1/config.json @@ -0,0 +1,27 @@ +{ + "cache_roots": "demucs_results_all", + "val_cache_roots": "demucs_results_val_full", + "out_dir": "restoflow_runs/gen_advramp_v1", + "target_stems": "other,vocals,drums,bass", + "context_sources": "restored,degraded", + "epochs": 50, + "batch": 256, + "lr": 0.001, + "hidden": 512, + "depth": 6, + "arch": "conv", + "heads": 8, + "sigma": 1.0, + "cfg_drop": 0.1, + "cfg_w": 2.0, + "cfg_rescale": 0.0, + "dist_loss": "none", + "dist_loss_w": 0.0, + "sample_steps": 40, + "eval_every": 25, + "eval_pairs": 150, + "demo_pairs": 8, + "device": "cuda", + "num_workers": 4, + "smoke": false +} \ No newline at end of file diff --git a/restoflow_runs/gen_advramp_v1/gen_stats.pt b/restoflow_runs/gen_advramp_v1/gen_stats.pt new file mode 100644 index 0000000000000000000000000000000000000000..cccdbb369772987b8ba05908900d2f2eb41de97a --- /dev/null +++ b/restoflow_runs/gen_advramp_v1/gen_stats.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:df1a99ec421f853c4e6165d41be3eda43d10cb86df3bf408d49a33c3c2a8bbbb +size 21727 diff --git a/restoflow_runs/gen_distvar_baseline/ckpt_best.pt b/restoflow_runs/gen_distvar_baseline/ckpt_best.pt new file mode 100644 index 0000000000000000000000000000000000000000..f35179ffa3b31b25b653810a50fe0eddf3d95274 --- /dev/null +++ b/restoflow_runs/gen_distvar_baseline/ckpt_best.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02f961a776b176809359dec9a33a5b322cc80628ae61c89a96ffa6c0f0591ee5 +size 41804914 diff --git a/restoflow_runs/gen_distvar_baseline/config.json b/restoflow_runs/gen_distvar_baseline/config.json new file mode 100644 index 0000000000000000000000000000000000000000..33dc8b98374bb2c142edd8339555c4b28bd29273 --- /dev/null +++ b/restoflow_runs/gen_distvar_baseline/config.json @@ -0,0 +1,27 @@ +{ + "cache_roots": "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_all", + "val_cache_roots": "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_val_full", + "out_dir": "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/restoflow_runs/gen_distvar_baseline", + "target_stems": "other,vocals,drums,bass", + "context_sources": "restored,degraded", + "epochs": 50, + "batch": 256, + "lr": 0.001, + "hidden": 512, + "depth": 6, + "arch": "conv", + "heads": 8, + "sigma": 1.0, + "cfg_drop": 0.1, + "cfg_w": 2.0, + "cfg_rescale": 0.0, + "dist_loss": "none", + "dist_loss_w": 0.0, + "sample_steps": 40, + "eval_every": 10, + "eval_pairs": 160, + "demo_pairs": 6, + "device": "cuda", + "num_workers": 4, + "smoke": false +} \ No newline at end of file diff --git a/restoflow_runs/gen_distvar_baseline/gen_stats.pt b/restoflow_runs/gen_distvar_baseline/gen_stats.pt new file mode 100644 index 0000000000000000000000000000000000000000..7a9ee28722f220153d6c0acf7cb77d717b4d4b7d --- /dev/null +++ b/restoflow_runs/gen_distvar_baseline/gen_stats.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:02d52658d03e36af2b24ccabb5ba9fd0e88b059df696551b743b4dd3bd2c62a2 +size 21727 diff --git a/restoflow_runs/restorer_attn_advramp/ckpt_best.pt b/restoflow_runs/restorer_attn_advramp/ckpt_best.pt new file mode 100644 index 0000000000000000000000000000000000000000..c8815ae753eaccefd3bf5d2b1a411d1ae6166c5d --- /dev/null +++ b/restoflow_runs/restorer_attn_advramp/ckpt_best.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0e39686522c2692f19e891105807421f7eebc8a7e6ecacf6387c7708867bab2e +size 134361163 diff --git a/restoflow_runs/restorer_attn_advramp/config.json b/restoflow_runs/restorer_attn_advramp/config.json new file mode 100644 index 0000000000000000000000000000000000000000..76f9eff6adb8108bef264ef5d51ee31554147df9 --- /dev/null +++ b/restoflow_runs/restorer_attn_advramp/config.json @@ -0,0 +1,45 @@ +{ + "cache_roots": [ + "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_all" + ], + "val_cache_roots": [ + "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_val_full" + ], + "out_dir": "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/restoflow_runs/restorer_attn_distvar_w003", + "stems": [ + "other", + "vocals", + "drums", + "bass" + ], + "T": 32, + "latent_dim": 256, + "val_frac": 0.1, + "split_seed": 0, + "rms_thr": -40.0, + "peak_thr": -25.0, + "max_drop_db": 15.0, + "stats_path": "", + "model_kind": "attn", + "use_mix": true, + "hidden": 512, + "depth": 8, + "stem_emb_dim": 32, + "heads": 8, + "sigma": 0.0, + "stem_weight_alpha": 0.5, + "energy_w": 0.0, + "dist_loss": "var", + "dist_loss_w": 0.003, + "batch": 256, + "lr": 0.001, + "wd": 0.0001, + "epochs": 50, + "device": "cuda", + "num_workers": 4, + "eval_every": 10, + "sample_steps": 40, + "eval_max_pairs": 400, + "demo_pairs": 6, + "model_id": "stabilityai/SAME-L" +} \ No newline at end of file diff --git a/restoflow_runs/restorer_attn_advramp/norm_stats.pt b/restoflow_runs/restorer_attn_advramp/norm_stats.pt new file mode 100644 index 0000000000000000000000000000000000000000..d13d88490dc821a84281eba1fb4bceed5170e8cf --- /dev/null +++ b/restoflow_runs/restorer_attn_advramp/norm_stats.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:17a91efd7e6743da59527409dd861f606c84ea1700ad2be2c183b0a0f3f58f41 +size 14197 diff --git a/restoflow_runs/restorer_attn_w003_gan/ckpt_best.pt b/restoflow_runs/restorer_attn_w003_gan/ckpt_best.pt new file mode 100644 index 0000000000000000000000000000000000000000..d0345cef0d8cf03b98f417172427975069e13faf --- /dev/null +++ b/restoflow_runs/restorer_attn_w003_gan/ckpt_best.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4e6f189fc339d593ab56034a66ef3532e251e18f1d93900833428aa9cf125cec +size 134360971 diff --git a/restoflow_runs/restorer_attn_w003_gan/config.json b/restoflow_runs/restorer_attn_w003_gan/config.json new file mode 100644 index 0000000000000000000000000000000000000000..76f9eff6adb8108bef264ef5d51ee31554147df9 --- /dev/null +++ b/restoflow_runs/restorer_attn_w003_gan/config.json @@ -0,0 +1,45 @@ +{ + "cache_roots": [ + "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_all" + ], + "val_cache_roots": [ + "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/demucs_results_val_full" + ], + "out_dir": "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra/restoflow_runs/restorer_attn_distvar_w003", + "stems": [ + "other", + "vocals", + "drums", + "bass" + ], + "T": 32, + "latent_dim": 256, + "val_frac": 0.1, + "split_seed": 0, + "rms_thr": -40.0, + "peak_thr": -25.0, + "max_drop_db": 15.0, + "stats_path": "", + "model_kind": "attn", + "use_mix": true, + "hidden": 512, + "depth": 8, + "stem_emb_dim": 32, + "heads": 8, + "sigma": 0.0, + "stem_weight_alpha": 0.5, + "energy_w": 0.0, + "dist_loss": "var", + "dist_loss_w": 0.003, + "batch": 256, + "lr": 0.001, + "wd": 0.0001, + "epochs": 50, + "device": "cuda", + "num_workers": 4, + "eval_every": 10, + "sample_steps": 40, + "eval_max_pairs": 400, + "demo_pairs": 6, + "model_id": "stabilityai/SAME-L" +} \ No newline at end of file diff --git a/restoflow_runs/restorer_attn_w003_gan/norm_stats.pt b/restoflow_runs/restorer_attn_w003_gan/norm_stats.pt new file mode 100644 index 0000000000000000000000000000000000000000..d13d88490dc821a84281eba1fb4bceed5170e8cf --- /dev/null +++ b/restoflow_runs/restorer_attn_w003_gan/norm_stats.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:17a91efd7e6743da59527409dd861f606c84ea1700ad2be2c183b0a0f3f58f41 +size 14197