{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, glob, copy, numpy as np, pandas as pd, librosa, matplotlib.pyplot as plt\nimport torch, torch.nn as nn, torch.optim as optim\nimport warnings; warnings.filterwarnings(\"ignore\"); np.random.seed(0); torch.manual_seed(0)\n\nSR=22050; N_FFT=1024; HOP=256; N_MELS=128\nSEG_SEC=1.0; SEG_FR=int(np.ceil(SEG_SEC*SR/HOP)); CLEN=int(SEG_SEC*SR); EMB=128\nFT_STEPS=150; FT_LR=1e-4; FT_BATCH=16; POS_THRESH=0.40   # low floor: capture candidates for sweeping\nOP=0.55                                                  # operating threshold (the deck's GATE)\nN_BEDS=20\ndevice=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"); print(\"device\",device)\ndef f2t(f): return f*HOP/SR\ndef rms(x): return float(np.sqrt(np.mean(x**2)+1e-12))\ndef load(p,dur=None): y,_=librosa.load(p,sr=SR,mono=True,duration=dur); return y.astype(\"float32\")\ndef pcen_wave(y):\n    S=librosa.feature.melspectrogram(y=y,sr=SR,n_fft=N_FFT,hop_length=HOP,n_mels=N_MELS,power=1.0)\n    return librosa.pcen(S*(2**31),sr=SR,hop_length=HOP).astype(\"float32\")\ndef _crop(feat,c):\n    a=c-SEG_FR//2; p=np.full((N_MELS,SEG_FR),feat.min(),dtype=\"float32\")\n    a0,a1=max(0,a),min(feat.shape[1],a+SEG_FR); p[:,(a0-a):(a0-a)+(a1-a0)]=feat[:,a0:a1]; return p\ndef pcen_window(y):\n    feat=pcen_wave(y)\n    if feat.shape[1]<=SEG_FR: return _crop(feat,feat.shape[1]//2)\n    col=feat.sum(0); cs=np.concatenate([[0],np.cumsum(col)]); win=cs[SEG_FR:]-cs[:-SEG_FR]\n    return _crop(feat,int(np.argmax(win))+SEG_FR//2)\ndef loudest_window(y,dur=SEG_SEC):\n    n=int(dur*SR)\n    if len(y)<=n: return np.pad(y,(0,max(0,n-len(y))))[:n]\n    e=y.astype(\"float64\")**2; cs=np.concatenate([[0],np.cumsum(e)]); w=cs[n:]-cs[:-n]\n    return y[int(np.argmax(w)):int(np.argmax(w))+n]\ndef top_windows(y,k=2,dur=SEG_SEC):\n    n=int(dur*SR)\n    if len(y)<=n: return [np.pad(y,(0,max(0,n-len(y))))[:n]]\n    e=y.astype(\"float64\")**2; cs=np.concatenate([[0],np.cumsum(e)]); w=cs[n:]-cs[:-n]\n    out=[];taken=[]\n    for _ in range(k):\n        pick=None\n        for idx in np.argsort(w)[::-1]:\n            if all(abs(idx-t)>n for t in taken): pick=int(idx);break\n        if pick is None: break\n        taken.append(pick); out.append(y[pick:pick+n])\n    return out\nclass Encoder(nn.Module):\n    def __init__(s,emb=EMB):\n        super().__init__()\n        def b(i,o): return nn.Sequential(nn.Conv2d(i,o,3,padding=1,bias=False),nn.BatchNorm2d(o),nn.ReLU(),nn.MaxPool2d(2))\n        s.enc=nn.Sequential(b(1,128),b(128,128),b(128,128),b(128,emb)); s.pool=nn.AdaptiveAvgPool2d(1)\n    def forward(s,x): return s.pool(s.enc(x.unsqueeze(1))).flatten(1)\ndef proto_logits(m,sx,sy,qx):\n    se,qe=m(sx),m(qx); pr=torch.stack([se[sy==c].mean(0) for c in torch.unique(sy)]); return -torch.cdist(qe,pr)\ndef spec_augment(x,F=20,T=10,nf=2,nt=2):\n    x=x.copy()\n    for _ in range(nf):\n        k=np.random.randint(0,F+1); a=np.random.randint(0,max(1,N_MELS-k)); x[a:a+k,:]=x.min()\n    for _ in range(nt):\n        k=np.random.randint(0,T+1); a=np.random.randint(0,max(1,SEG_FR-k)); x[:,a:a+k]=x.min()\n    return x\ndef augment(x): return spec_augment(x) if np.random.rand()<0.6 else x\ndef train_encoder(bank,episodes=1000,n_way=6,k=4,q=4):\n    m=Encoder().to(device); opt=optim.Adam(m.parameters(),1e-3); ce=nn.CrossEntropyLoss(); cls=list(bank); m.train()\n    for ep in range(episodes):\n        cc=list(np.random.choice(cls,min(n_way,len(cls)),replace=False)); sx,sy,qx,qy=[],[],[],[]\n        for l,c in enumerate(cc):\n            idx=np.random.permutation(len(bank[c]))[:k+q]\n            for j,i in enumerate(idx):\n                (sx if j<k else qx).append(torch.tensor(augment(bank[c][i]))); (sy if j<k else qy).append(l)\n        sx=torch.stack(sx).to(device);qx=torch.stack(qx).to(device);sy=torch.tensor(sy).to(device);qy=torch.tensor(qy).to(device)\n        loss=ce(proto_logits(m,sx,sy,qx),qy); opt.zero_grad(); loss.backward(); opt.step()\n        if (ep+1)%250==0: print(\"  enc ep %d loss %.3f\"%(ep+1,loss.item()))\n    return m\ndef mix(bed, calls, n, snr_db, gap=0.4):\n    y=bed.copy().astype(\"float32\"); L=len(y); ivs=[]; tries=0\n    while len(ivs)<n and tries<n*60:\n        tries+=1; c=calls[np.random.randint(len(calls))][:CLEN]; c=np.pad(c,(0,max(0,CLEN-len(c)))).astype(\"float32\")\n        if L<=CLEN: break\n        pos=np.random.randint(0,L-CLEN)\n        if any(pos<e+int(gap*SR) and pos+CLEN>s-int(gap*SR) for s,e in ivs): continue\n        snr=np.random.uniform(*snr_db) if isinstance(snr_db,tuple) else snr_db\n        sc=rms(y[pos:pos+CLEN])*(10**(snr/20))/(rms(c)+1e-9); y[pos:pos+CLEN]+=c*sc; ivs.append((pos,pos+CLEN))\n    return y, sorted(ivs)\ndef detect_scored(base, bed_wave, support_waves):\n    \"\"\"Return candidate events with their peak probability (no gate cut), plus frame probs.\"\"\"\n    train_wave, iv = mix(bed_wave, support_waves, 45, (-6,12))\n    tf=pcen_wave(train_wave); sf=pcen_wave(bed_wave)\n    pf=[((s+e)//2)//HOP for s,e in iv]; pos=[_crop(tf,f) for f in pf]\n    negs=[]\n    while len(negs)<150:\n        f=np.random.randint(SEG_FR, tf.shape[1]-SEG_FR)\n        if all(abs(f-p)>SEG_FR for p in pf): negs.append(_crop(tf,f))\n    m=copy.deepcopy(base).to(device); h=nn.Linear(EMB,2).to(device)\n    opt=optim.Adam(list(m.parameters())+list(h.parameters()),FT_LR); ce=nn.CrossEntropyLoss(); m.train();h.train()\n    for _ in range(FT_STEPS):\n        xs,ys=[],[]\n        for _ in range(FT_BATCH//2):\n            xs.append(augment(pos[np.random.randint(len(pos))])); ys.append(1)\n            xs.append(negs[np.random.randint(len(negs))]); ys.append(0)\n        xb=torch.tensor(np.stack(xs)).to(device); yb=torch.tensor(ys).to(device)\n        loss=ce(h(m(xb)),yb); opt.zero_grad(); loss.backward(); opt.step()\n    m.eval();h.eval(); st=max(1,SEG_FR//3); frames=list(range(SEG_FR,sf.shape[1]-SEG_FR,st)); probs=[]\n    with torch.no_grad():\n        for i in range(0,len(frames),256):\n            x=np.stack([_crop(sf,f) for f in frames[i:i+256]])\n            probs.extend(torch.softmax(h(m(torch.tensor(x).to(device))),1)[:,1].cpu().numpy().tolist())\n    probs=np.asarray(probs,dtype=\"float32\"); above=probs>=POS_THRESH; ev=[]; i=0\n    while i<len(frames):\n        if above[i]:\n            j=i\n            while j+1<len(frames) and above[j+1] and frames[j+1]-frames[j]<=st*1.5: j+=1\n            ev.append((max(0,f2t(frames[i])-SEG_SEC/2), f2t(frames[j])+SEG_SEC/2, float(probs[i:j+1].max())))\n            i=j+1\n        else: i+=1\n    return ev, frames, probs\n\ndef match(dets, gts):\n    order=sorted(range(len(dets)),key=lambda i:-dets[i][2]); matched={}; scored=[None]*len(dets)\n    for i in order:\n        s,e,p=dets[i]; hit=False\n        for gi,(gs,ge) in enumerate(gts):\n            if gi in matched: continue\n            if max(0,min(e,ge)-max(s,gs))>0: matched[gi]=p; hit=True; break\n        scored[i]=(p,hit)\n    gtp=[matched.get(gi,0.0) for gi in range(len(gts))]\n    return scored, gtp\ndef pr_at(scored,total,thr):\n    tp=sum(1 for p,m in scored if p>=thr and m); fp=sum(1 for p,m in scored if p>=thr and not m); fn=total-tp\n    P=tp/(tp+fp) if tp+fp else 1.0; R=tp/total if total else 0.0\n    return P,R,(2*P*R/(P+R) if P+R else 0.0),tp,fp,fn\n\n# ---- train + evaluate ----\nROOT=os.path.dirname(glob.glob('/kaggle/input/**/train_metadata.csv',recursive=True)[0])\nAUD=os.path.join(ROOT,'train_audio'); SS=os.path.join(ROOT,'unlabeled_soundscapes'); m=pd.read_csv(os.path.join(ROOT,'train_metadata.csv'))\ntg=m[m.primary_label=='whbsho3'].sort_values('rating',ascending=False); fns=tg['filename'].tolist()\nsupport=[loudest_window(load(os.path.join(AUD,fn),dur=30)) for fn in fns[:5]]\nhide=[]\nfor fn in fns[5:]:\n    try: hide+=top_windows(load(os.path.join(AUD,fn),dur=30),2)\n    except: pass\nothers=[s for s in m.primary_label.unique() if s!='whbsho3']; np.random.shuffle(others); bank={}\nfor sp in others:\n    if len(bank)>=40: break\n    segs=[]\n    for fn in m[m.primary_label==sp]['filename'].tolist()[:12]:\n        try: segs.append(pcen_window(load(os.path.join(AUD,fn),dur=15)))\n        except: pass\n    if len(segs)>=8: bank[sp]=segs\nprint(\"training encoder...\"); model=train_encoder(bank,episodes=1000)\n\nsnr_levels=[9,6,3,0,-3,-6]; all_scored=[]; total_gt=0; per_call=[]; facc=[0,0,0,0]  # tp,tn,fp,fn frames\nbed_paths=sorted(glob.glob(SS+'/*.ogg'))[:N_BEDS]\nprint(\"evaluating on %d soundscapes...\"%len(bed_paths))\nfor bi,bp in enumerate(bed_paths):\n    bed=load(bp,dur=240); y=bed.copy(); L=len(y); used=[]; ivs=[]; snrs=[]\n    for kk in range(12):\n        s=snr_levels[kk%len(snr_levels)]; c=hide[np.random.randint(len(hide))][:CLEN]; c=np.pad(c,(0,max(0,CLEN-len(c)))).astype(\"float32\")\n        ok=False\n        for _ in range(80):\n            pos=np.random.randint(0,L-CLEN)\n            if any(pos<e+int(0.4*SR) and pos+CLEN>a-int(0.4*SR) for a,e in used): continue\n            ok=True; break\n        if not ok: continue\n        sc=rms(y[pos:pos+CLEN])*(10**(s/20))/(rms(c)+1e-9); y[pos:pos+CLEN]+=c*sc\n        used.append((pos,pos+CLEN)); ivs.append((pos/SR,(pos+CLEN)/SR)); snrs.append(s)\n    dets,frames,probs=detect_scored(model,y,support)\n    scored,gtp=match([(s,e,p) for s,e,p in dets], ivs)\n    all_scored+=scored; total_gt+=len(ivs)\n    for gi in range(len(ivs)): per_call.append((snrs[gi], gtp[gi]))\n    for f,p in zip(frames,probs):\n        ts=f2t(f); iscall=any(s<=ts<=e for s,e in ivs); pred=p>=OP\n        if iscall and pred: facc[0]+=1\n        elif (not iscall) and (not pred): facc[1]+=1\n        elif (not iscall) and pred: facc[2]+=1\n        else: facc[3]+=1\n    if (bi+1)%5==0: print(\"  %d/%d beds\"%(bi+1,len(bed_paths)))\n\n# ---- report ----\nthrs=np.round(np.arange(0.30,0.96,0.02),2); curve=[pr_at(all_scored,total_gt,t) for t in thrs]\nP,R,F1,TP,FP,FN=pr_at(all_scored,total_gt,OP)\nbi=int(np.argmax([c[2] for c in curve])); bP,bR,bF1,*_=curve[bi]\n# average precision = area under P-R (integrate P over R)\nRs=np.array([c[1] for c in curve]); Ps=np.array([c[0] for c in curve]); o=np.argsort(Rs); AP=float(np.trapz(Ps[o],Rs[o]))\nacc=(facc[0]+facc[1])/sum(facc)\nprint(\"\\n============ ARGUS DETECTOR EVALUATION ============\")\nprint(\"planted calls: %d over %d soundscapes\"%(total_gt,len(bed_paths)))\nprint(\"\\n--- at operating point (thr=%.2f) ---\"%OP)\nprint(\"  Precision = %.1f%%   Recall = %.1f%%   F1 = %.3f   (TP=%d FP=%d FN=%d)\"%(100*P,100*R,F1,TP,FP,FN))\nprint(\"  best-F1 point: thr=%.2f  P=%.1f%%  R=%.1f%%  F1=%.3f\"%(thrs[bi],100*bP,100*bR,bF1))\nprint(\"  Average Precision (area under P-R) = %.3f\"%AP)\nprint(\"\\n--- recall by call loudness (at thr=%.2f) ---\"%OP)\nfor s in snr_levels:\n    v=[gp for (ss,gp) in per_call if ss==s]; rec=np.mean([g>=OP for g in v]) if v else 0\n    print(\"  %+d dB : recall %.0f%% (%d calls)\"%(s,100*rec,len(v)))\nprint(\"\\n--- window-level accuracy = %.1f%% ---  (NOTE: misleading — background dominates)\"%(100*acc))\n\n# PR curve figure\nNAVY=\"#0E1525\";EM=\"#34D399\";MUT=\"#9FB0C9\"\nfig,ax=plt.subplots(figsize=(4,3),dpi=220); fig.patch.set_facecolor(NAVY); ax.set_facecolor(NAVY)\nax.plot([c[1] for c in curve],[c[0] for c in curve],\"-o\",color=EM,lw=2,ms=3)\nax.scatter([R],[P],color=\"#FFC861\",zorder=5,s=40,label=\"operating point\")\nax.set_xlabel(\"recall\",color=MUT); ax.set_ylabel(\"precision\",color=MUT); ax.set_xlim(0,1); ax.set_ylim(0,1.05)\nax.tick_params(colors=MUT)\nfor sp in ax.spines.values(): sp.set_color(\"#2C3A5C\")\nax.grid(True,color=\"#1c2a44\",lw=0.6); ax.legend(fontsize=7,labelcolor=MUT,frameon=False)\nax.set_title(\"ARGUS precision–recall (AP=%.2f)\"%AP,color=\"#F2F6FC\",fontsize=9)\nfig.savefig(\"/kaggle/working/argus_pr_curve.png\",facecolor=NAVY,bbox_inches=\"tight\"); print(\"\\nsaved argus_pr_curve.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-23T15:57:20.130071Z","iopub.execute_input":"2026-07-23T15:57:20.130391Z","iopub.status.idle":"2026-07-23T16:02:06.711196Z","shell.execute_reply.started":"2026-07-23T15:57:20.130342Z","shell.execute_reply":"2026-07-23T16:02:06.710273Z"}},"outputs":[],"execution_count":null}]}