ganesh333's picture
Upload folder using huggingface_hub
001425d verified
Raw History Blame Contribute Delete
2.89 kB
import unittest, json
from pathlib import Path
import numpy as np, torch
from src.model import InsectCNN
from src.inference import load_model
from src.inference import identify, genus_of
from src.audio import logmel, crop_offsets, resample, SR
ROOT=Path(__file__).resolve().parents[1]
class ModelPackageTests(unittest.TestCase):
def test_forward_shape(self):
m=InsectCNN(32).eval()
with torch.no_grad(): out=m(torch.zeros(2,1,64,201))
self.assertEqual(tuple(out.shape),(2,32))
def test_hdf5_load(self):
m,labels,cfg=load_model(ROOT)
self.assertEqual(len(labels),32)
self.assertEqual(cfg['num_classes'],32)
def test_center_crop_default_unchanged(self):
y=np.random.default_rng(0).standard_normal(44100*5).astype(np.float32)
y16=resample(y,44100)
np.testing.assert_allclose(logmel(y,44100),logmel(y16,SR,(len(y16)-2*SR)//2),atol=1e-5)
self.assertEqual(crop_offsets(len(y16),1),[None])
offs=crop_offsets(len(y16),5)
self.assertEqual((offs[0],offs[-1]),(0,len(y16)-2*SR))
def test_finetuned_bundle_loads(self):
d=ROOT/'finetuned_cnn'
if not (d/'model.h5').exists(): self.skipTest('finetuned_cnn not trained')
m,labels,cfg=load_model(d)
self.assertEqual(len(labels),32)
def test_identify_levels(self):
labels={str(i):n for i,n in enumerate(['Myopsaltaleona','Myopsaltamackinlayi','Chorthippusbrunneus','Pseudochorthippusparallelus'])}
self.assertEqual(genus_of('Pseudochorthippusparallelus'),'Pseudochorthippus')
self.assertEqual(identify(np.array([.7,.1,.1,.1]),labels)['level'],'species')
r=identify(np.array([.4,.35,.15,.1]),labels); self.assertEqual((r['level'],r['name']),('genus','Myopsalta'))
self.assertEqual(identify(np.array([.15,.15,.35,.35]),labels,0.6)['name'],'Orthoptera (crickets/grasshoppers)')
self.assertEqual(identify(np.array([.3,.2,.3,.2]),labels,0.6)['level'],'unknown')
def test_litert_parity(self):
tflite_path = ROOT / 'model.tflite'
if not tflite_path.exists():
self.skipTest('model.tflite not exported')
import tensorflow as tf
interpreter = tf.lite.Interpreter(model_path=str(tflite_path))
interpreter.allocate_tensors()
in_idx = interpreter.get_input_details()[0]['index']
out_idx = interpreter.get_output_details()[0]['index']
m, labels, cfg = load_model(ROOT)
sample = np.random.default_rng(123).standard_normal((1, 1, 64, 201)).astype(np.float32)
with torch.no_grad():
pt_out = m(torch.from_numpy(sample)).cpu().numpy()
interpreter.set_tensor(in_idx, sample)
interpreter.invoke()
tflite_out = interpreter.get_tensor(out_idx)
np.testing.assert_allclose(pt_out, tflite_out, atol=1e-4)
if __name__=='__main__': unittest.main()