import sys,tempfile,unittest from pathlib import Path import numpy as np from PIL import Image ROOT=Path(__file__).resolve().parents[1] sys.path[:0]=[str(ROOT/'code'),str(ROOT/'ui')] from train_recon import load_dat,Net from infer_brainmu_lora import metrics,LoRALinear from multi_gpu import shard_indices import torch class CoreTests(unittest.TestCase): def test_dat_orientation(self): x=np.zeros((41,250,400),np.uint8);x[0,0,0]=1;x[-1,-1,-1]=1 with tempfile.TemporaryDirectory() as d: p=Path(d)/'x.dat';np.packbits(x.reshape(41,-1),axis=1,bitorder='little').tofile(p) np.testing.assert_array_equal(load_dat(p),x[:,::-1,:]) def test_metrics_identity(self): m=metrics(Image.new('RGB',(32,32),'gray'),Image.new('RGB',(32,32),'gray')) self.assertAlmostEqual(m['ssim'],1);self.assertAlmostEqual(m['psnr_db'],120) def test_shards(self): s=shard_indices(1000,8);self.assertEqual([len(x) for x in s],[125]*8) self.assertEqual(sorted(i for x in s for i in x),list(range(1000))) def test_lora_zero(self): base=torch.nn.Linear(4,3);m=LoRALinear(base,2,4,0);m.lora_A.data.zero_();m.lora_B.data.zero_();x=torch.ones(1,4) torch.testing.assert_close(m(x),base(x)) def test_frontend(self): self.assertEqual(tuple(Net()(torch.zeros(1,41,8,8)).shape),(1,1,8,8)) if __name__=='__main__':unittest.main()