Allow multi-section YuE2 lyrics and force non-flash attention on Windows.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
spec = importlib.util.spec_from_file_location('yue2_worker', Path(__file__).parents[1] / 'scripts/yue2-worker.py')
|
||||
worker = importlib.util.module_from_spec(spec)
|
||||
@@ -9,12 +10,29 @@ spec.loader.exec_module(worker)
|
||||
|
||||
|
||||
class WorkerTests(unittest.TestCase):
|
||||
def test_lyrics_preserve_words_and_normalize_ui_headings(self):
|
||||
self.assertEqual(worker.normalize_lyrics('[Pre-Chorus]\nEvery word stays\n[Outro]\n'), '[prechorus]\nEvery word stays\n\n')
|
||||
self.assertEqual(worker.normalize_lyrics('[Verse 1]\nHello'), '[verse1]\nHello\n\n')
|
||||
for lyrics in ['', 'No heading', '[Verse]\nA\n[Chorus]\nB']:
|
||||
with self.assertRaises(ValueError):
|
||||
worker.normalize_lyrics(lyrics)
|
||||
def test_lyrics_allow_multiple_sections_and_wrap_plain_text(self):
|
||||
multi = '[Verse 1]\nHello\n\n[Chorus]\nSing it'
|
||||
self.assertEqual(worker.normalize_lyrics(multi), multi + '\n')
|
||||
self.assertEqual(worker.normalize_lyrics('[Verse 1]\nHello'), '[Verse 1]\nHello\n')
|
||||
self.assertEqual(worker.normalize_lyrics('No heading yet'), '[song]\nNo heading yet\n\n')
|
||||
with self.assertRaises(ValueError):
|
||||
worker.normalize_lyrics('')
|
||||
with self.assertRaises(ValueError):
|
||||
worker.normalize_lyrics(' ')
|
||||
|
||||
def test_attention_backend_never_selects_flash(self):
|
||||
torch_mod = SimpleNamespace(
|
||||
cuda=SimpleNamespace(is_available=lambda: True),
|
||||
backends=SimpleNamespace(cudnn=SimpleNamespace(is_available=lambda: True)),
|
||||
)
|
||||
attention, pipeline = worker.resolve_attention_backend(torch_mod)
|
||||
self.assertEqual(attention, 'cudnn')
|
||||
self.assertEqual(pipeline, 'torch')
|
||||
self.assertNotEqual(attention, 'flash')
|
||||
torch_mod.backends.cudnn.is_available = lambda: False
|
||||
attention, pipeline = worker.resolve_attention_backend(torch_mod)
|
||||
self.assertEqual(attention, 'sdpa')
|
||||
self.assertEqual(pipeline, 'torch-eager')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
Reference in New Issue
Block a user