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 train.py from ganesh333/kernel-insect-classifier: direct link, hf CLI and curl.
- Browser
- Download file 9.9 kB
-
https://huggingface.co/ganesh333/kernel-insect-classifier/resolve/main/train.py
- Command line
-
hf download hf://ganesh333/kernel-insect-classifier/train.py
-
curl -L -o train.py https://huggingface.co/ganesh333/kernel-insect-classifier/resolve/main/train.py
9.9 kB
| import os, io, re, json, zipfile, argparse, random | |
| from pathlib import Path | |
| from urllib.parse import unquote | |
| import numpy as np, pandas as pd, torch, h5py | |
| from scipy.io import wavfile | |
| from sklearn.metrics import accuracy_score, classification_report, f1_score | |
| from src.audio import logmel, resample, crop_offsets, SR, DURATION | |
| from src.model import InsectCNN | |
| ROOT=Path(__file__).resolve().parent | |
| def load_rows(data_dir=ROOT/'data'): | |
| """Return dicts with the 16 kHz waveform, label, split and source session of every CSV row found in the zips.""" | |
| rows=[] | |
| for csv,zipname in [('Cicadidae.csv','Cicadidae.zip'),('Orthoptera.csv','Orthoptera.zip')]: | |
| df=pd.read_csv(Path(data_dir)/csv) | |
| with zipfile.ZipFile(Path(data_dir)/zipname) as z: | |
| wavs=[n for n in z.namelist() if n.lower().endswith('.wav') and '/._' not in n] | |
| lookup={} | |
| for n in wavs: lookup.setdefault(n.split('/')[-1],n) | |
| for _,r in df.iterrows(): | |
| fn=str(r['file_name']) | |
| member=lookup.get(fn) | |
| if not member: | |
| matches=[n for n in wavs if n.endswith('/'+fn)] | |
| if matches: member=matches[0] | |
| if not member: | |
| print('MISSING',fn); continue | |
| sr,y=wavfile.read(io.BytesIO(z.read(member))) | |
| if y.ndim==2: y=y.astype(np.float32).mean(axis=1) | |
| if np.issubdtype(y.dtype,np.integer): | |
| y=y.astype(np.float32)/max(abs(np.iinfo(y.dtype).min),np.iinfo(y.dtype).max) | |
| else: y=y.astype(np.float32) | |
| rows.append({'wave':resample(y,sr),'label':int(r['class_ID']),'species':str(r['species']), | |
| 'split':str(r['data_set']).lower(),'file_name':fn, | |
| 'session':session_key(str(r['species']),str(r['original_file_name']))}) | |
| return rows | |
| def session_key(species, original_file_name): | |
| """Group files cut from the same source recording/tape (e.g. 'MHV 313 ... #6a/#6b', 'dat008-019').""" | |
| s=unquote(original_file_name) | |
| for p in [r'(MHV\s*\d+)',r'(dat\d+)',r'([A-Z]+_\d{8})',r'(\d{6})_']: | |
| m=re.match(p,s) | |
| if m: return species+'|'+m.group(1) | |
| return species+'|'+s.rsplit('.',1)[0] | |
| def features(waves, offsets): | |
| return torch.from_numpy(np.stack([logmel(w,SR,o) for w,o in zip(waves,offsets)])) | |
| def predict_probs(model, rows, idx, eval_crops=1): | |
| """Per-file class probabilities, averaged over eval_crops evenly spaced 2 s windows.""" | |
| model.eval(); out=[] | |
| with torch.no_grad(): | |
| for i in idx: | |
| offs=crop_offsets(len(rows[i]['wave']),eval_crops) | |
| out.append(torch.softmax(model(features([rows[i]['wave']]*len(offs),offs)),dim=1).mean(dim=0).numpy()) | |
| return np.stack(out) if out else np.empty((0,32)) | |
| def evaluate_split(model, rows, idx, eval_crops=1, names=None, n_boot=1000, seed=42): | |
| y=np.array([rows[i]['label'] for i in idx]); pred=np.argmax(predict_probs(model,rows,idx,eval_crops),axis=1) | |
| res={'count':int(len(idx)),'accuracy':float(accuracy_score(y,pred)), | |
| 'macro_f1':float(f1_score(y,pred,average='macro',zero_division=0)), | |
| 'weighted_f1':float(f1_score(y,pred,average='weighted',zero_division=0))} | |
| # Bootstrap over files: with ~74 test files the point estimates are very uncertain. | |
| rng=np.random.default_rng(seed); acc=[]; mf1=[] | |
| for _ in range(n_boot): | |
| b=rng.integers(0,len(y),len(y)) | |
| acc.append(accuracy_score(y[b],pred[b])); mf1.append(f1_score(y[b],pred[b],average='macro',zero_division=0)) | |
| res['accuracy_ci95']=[float(np.percentile(acc,2.5)),float(np.percentile(acc,97.5))] | |
| res['macro_f1_ci95']=[float(np.percentile(mf1,2.5)),float(np.percentile(mf1,97.5))] | |
| if names: | |
| res['report']=classification_report(y,pred,labels=list(range(32)),target_names=[names[str(i)] for i in range(32)],output_dict=True,zero_division=0) | |
| return res | |
| def evaluate_test(model, rows, eval_crops=1, names=None): | |
| """Official test split, plus the subset whose source session never occurs in train/validation.""" | |
| split=np.array([r['split'] for r in rows]); test=np.where(split=='test')[0] | |
| seen={r['session'] for r in rows if r['split']!='test'} | |
| clean=np.array([i for i in test if rows[i]['session'] not in seen]) | |
| return {'test':evaluate_split(model,rows,test,eval_crops,names), | |
| 'test_session_clean':evaluate_split(model,rows,clean,eval_crops) if len(clean) else None} | |
| def main(): | |
| ap=argparse.ArgumentParser() | |
| ap.add_argument('--epochs',type=int,default=16); ap.add_argument('--seed',type=int,default=42) | |
| ap.add_argument('--output',default='.'); ap.add_argument('--data-dir',default=str(ROOT/'data')) | |
| ap.add_argument('--init-from',default=None,help='model dir to warm-start from (fine-tuning)') | |
| ap.add_argument('--lr',type=float,default=1e-3) | |
| ap.add_argument('--random-crop',action='store_true',help='train on a fresh random 2 s window per file each epoch') | |
| ap.add_argument('--eval-crops',type=int,default=1,help='average over K evenly spaced windows for validation/test') | |
| args=ap.parse_args() | |
| random.seed(args.seed); np.random.seed(args.seed); torch.manual_seed(args.seed) | |
| torch.set_num_threads(max(1,min(4,os.cpu_count() or 1))) | |
| rows=load_rows(args.data_dir) | |
| if not rows: raise RuntimeError('No audio rows loaded. Check data/ zip and CSV files.') | |
| y=np.array([r['label'] for r in rows]); split=np.array([r['split'] for r in rows]) | |
| train=np.where(split=='train')[0]; val=np.where(split=='validation')[0]; test=np.where(split=='test')[0] | |
| # Dataset-provided partitions are preserved; no random file-level split. Test is only scored once, after selection. | |
| waves=[r['wave'] for r in rows]; yt=torch.tensor(y,dtype=torch.long) | |
| if args.init_from: | |
| from src.inference import load_model | |
| model,_,_=load_model(args.init_from) | |
| else: | |
| model=InsectCNN(32) | |
| n=int(SR*DURATION) | |
| xcenter=features(waves,[None]*len(waves)) if not args.random_crop else None | |
| counts=np.bincount(y[train],minlength=32) | |
| weights=np.array([len(train)/(32*max(1,c)) for c in counts],dtype=np.float32) | |
| criterion=torch.nn.CrossEntropyLoss(weight=torch.tensor(weights)) | |
| opt=torch.optim.AdamW(model.parameters(),lr=args.lr,weight_decay=1e-4) | |
| best=-1; best_state=None; history=[] | |
| for epoch in range(args.epochs): | |
| model.train(); perm=np.random.permutation(train); losses=[] | |
| for start in range(0,len(perm),16): | |
| idx=perm[start:start+16] | |
| if args.random_crop: | |
| offs=[np.random.randint(0,len(waves[i])-n+1) if len(waves[i])>n else None for i in idx] | |
| xb=features([waves[i] for i in idx],offs) | |
| else: xb=xcenter[idx] | |
| opt.zero_grad(); logits=model(xb); loss=criterion(logits,yt[idx]); loss.backward(); opt.step(); losses.append(float(loss.item())) | |
| pv=np.argmax(predict_probs(model,rows,val,args.eval_crops),axis=1) if len(val) else np.array([]) | |
| score=f1_score(y[val],pv,average='macro',zero_division=0) if len(val) else -np.mean(losses) | |
| history.append({'epoch':epoch+1,'train_loss':float(np.mean(losses)),'validation_macro_f1':float(score)}) | |
| print(f"epoch {epoch+1}/{args.epochs}: loss={np.mean(losses):.4f} val_macro_f1={score:.4f}",flush=True) | |
| if score>best: | |
| best=score; best_state={k:v.detach().cpu().clone() for k,v in model.state_dict().items()} | |
| model.load_state_dict(best_state); model.eval() | |
| ids=sorted(set(int(i) for i in y)) | |
| names={str(i):next((r['species'] for r in rows if r['label']==i),f'class_{i}') for i in ids} | |
| ev=evaluate_test(model,rows,args.eval_crops,names); t=ev['test'] | |
| metrics={'dataset_files_loaded':len(rows),'split_counts':{s:int((split==s).sum()) for s in ['train','validation','test']}, | |
| 'num_classes':32,'test_accuracy':t['accuracy'],'test_macro_f1':t['macro_f1'],'test_weighted_f1':t['weighted_f1'], | |
| 'test_accuracy_ci95':t['accuracy_ci95'],'test_macro_f1_ci95':t['macro_f1_ci95'], | |
| 'test_report':t['report'],'test_session_clean':ev['test_session_clean'], | |
| 'best_validation_macro_f1':float(best),'eval_crops':args.eval_crops,'random_crop':args.random_crop, | |
| 'init_from':args.init_from,'lr':args.lr,'epochs':args.epochs, | |
| 'history':history,'seed':args.seed,'limitations':'Small imbalanced research dataset; results are not evidence of grain-pest detection.'} | |
| out=Path(args.output); out.mkdir(parents=True,exist_ok=True) | |
| torch.save(model.state_dict(),out/'model.pt') | |
| with h5py.File(out/'model.h5','w') as h: | |
| h.attrs['format']='kernel-pytorch-state-dict-hdf5-v1' | |
| h.attrs['model_architecture']='InsectCNN' | |
| h.attrs['model_version']='0.1.0-experimental' | |
| sd=h.create_group('state_dict') | |
| for k,v in model.state_dict().items(): sd.create_dataset(k,data=v.detach().cpu().numpy()) | |
| (out/'labels.json').write_text(json.dumps(names,indent=2)) | |
| (out/'config.json').write_text(json.dumps({'model_type':'kernel-insect-cnn','architecture':'InsectCNN','num_classes':32,'sample_rate':16000,'clip_seconds':2.0,'feature':'64-bin log-mel spectrogram','eval_crops':args.eval_crops,'model_version':'0.1.0-experimental','framework':'PyTorch','hdf5_format':'kernel-pytorch-state-dict-hdf5-v1','task':'insect species classification','experimental':True},indent=2)) | |
| (out/'metrics.json').write_text(json.dumps(metrics,indent=2)) | |
| (out/'history.json').write_text(json.dumps(history,indent=2)) | |
| print(json.dumps({k:v for k,v in metrics.items() if k not in ('test_report','history')},indent=2)) | |
| if __name__=='__main__': main() | |