Audio Classification
PyTorch
LiteRT
LiteRT
kernel-insect-cnn
insect
bioacoustics
experimental
local-inference
Instructions to use ganesh333/kernel-insect-classifier with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- LiteRT
How to use ganesh333/kernel-insect-classifier with LiteRT:
# No code snippets available yet for this library. # To use this model, check the repository files and the library's documentation. # Want to help? PRs adding snippets are welcome at: # https://github.com/huggingface/huggingface.js
- Notebooks
- Google Colab
- Kaggle
Download tests/test_model.py from ganesh333/kernel-insect-classifier: direct link, hf CLI and curl.
- Browser
- Download file 2.89 kB
-
https://huggingface.co/ganesh333/kernel-insect-classifier/resolve/main/tests/test_model.py
- Command line
-
hf download hf://ganesh333/kernel-insect-classifier/tests/test_model.py
-
curl -L -o test_model.py https://huggingface.co/ganesh333/kernel-insect-classifier/resolve/main/tests/test_model.py
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() | |