40 lines
2.5 KiB
Python
40 lines
2.5 KiB
Python
import tempfile
|
|
import unittest
|
|
from collections import namedtuple
|
|
from contextlib import nullcontext
|
|
from pathlib import Path
|
|
from unittest.mock import Mock,patch
|
|
from profiles import Profiles,SCHEMAS,video_parameters
|
|
from video import video_devices
|
|
from video_worker import EncoderOnDevice
|
|
|
|
class VideoDeviceTests(unittest.TestCase):
|
|
def test_selection_is_by_uuid_not_physical_index(self):
|
|
cards=[dict(uuid='GPU-second',name='RTX 3060',total_mib=12000),dict(uuid='GPU-first',name='RTX 5080',total_mib=16000)]
|
|
main,encoder,ordered=video_devices({'parameters':{'text_encoder_device':'GPU-second'}},cards)
|
|
self.assertEqual([g['uuid'] for g in ordered],['GPU-first','GPU-second'])
|
|
self.assertEqual(video_devices({'parameters':{}},cards)[2],[main])
|
|
self.assertEqual(video_devices({'parameters':{'video_device':'GPU-second'}},cards)[2],[encoder])
|
|
with self.assertRaises(ValueError):video_devices({'parameters':{'text_encoder_device':'GPU-missing'}},cards)
|
|
def test_parameters_persist_and_old_profiles_default_to_same_gpu(self):
|
|
with tempfile.TemporaryDirectory() as d:
|
|
catalog=Mock();catalog.entry.return_value=dict(id='m',kind='video',repo='example/repo',file='model.safetensors',profile_eligible=True)
|
|
store=Profiles(Path(d)/'profiles.json',catalog)
|
|
params={k:v[2] for k,v in SCHEMAS['video'].items()}
|
|
row=store.save(dict(id=None,revision=0,name='video-test',kind='video',model_id='m',parameters=params))
|
|
self.assertEqual(row['parameters']['text_encoder_device'],'same')
|
|
row['parameters'].update(video_device='GPU-first',text_encoder_device='GPU-second')
|
|
store.save({k:row[k] for k in ('id','revision','name','kind','model_id','parameters')})
|
|
self.assertEqual(Profiles(Path(d)/'profiles.json',catalog).rows[0]['parameters']['text_encoder_device'],'GPU-second')
|
|
for value in ['cpu','cuda:1',None,{},'GPU-abc;bad']:
|
|
with self.assertRaises(ValueError):video_parameters({'text_encoder_device':value})
|
|
def test_encoder_moves_conditioning_back_and_preserves_optional_audio(self):
|
|
result=namedtuple('Embedding','video_encoding audio_encoding attention_mask')
|
|
video,mask=Mock(),Mock();encoder=Mock(return_value=[result(video,None,mask)])
|
|
torch=Mock();torch.cuda.device.side_effect=lambda device:nullcontext()
|
|
with patch.dict('sys.modules',{'torch':torch}):
|
|
output=EncoderOnDevice(encoder,'cuda:1','cuda:0')(['synthetic'])
|
|
torch.cuda.device.assert_called_once_with('cuda:1')
|
|
video.to.assert_called_once_with('cuda:0');mask.to.assert_called_once_with('cuda:0')
|
|
self.assertIsNone(output[0].audio_encoding)
|