ganesh333's picture
Upload folder using huggingface_hub
575e16b verified
Raw History Blame Contribute Delete
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()