retrina0678's picture
Add model card / clean paths
db48d32 verified
|
Raw
History Blame Contribute Delete
7.04 kB
metadata
license: other
license_name: non-commercial-inherited
language:
  - ko
library_name: pytorch
pipeline_tag: audio-classification
tags:
  - audio
  - audio-classification
  - sound-event-detection
  - efficientat
  - mobilenetv3
  - edge
datasets:
  - retrina0678/miimo-audio-dataset
metrics:
  - accuracy
  - f1
  - recall
model-index:
  - name: miimo-efficientat-target4
    results:
      - task:
          type: audio-classification
          name: Sound Event Classification
        dataset:
          name: miimo-audio-dataset (4-class, held-out test)
          type: retrina0678/miimo-audio-dataset
        metrics:
          - type: accuracy
            value: 0.9774
            name: Accuracy (5-fold mean)
          - type: recall
            value: 0.9763
            name: Macro Recall (5-fold mean)
          - type: f1
            value: 0.9756
            name: Macro F1 (5-fold mean)

miimo-efficientat-target4

EfficientAT mn10_as ๋ฅผ 4๊ฐœ ์Œํ–ฅ ์ด๋ฒคํŠธ๋กœ ํŒŒ์ธํŠœ๋‹ํ•œ ๋ชจ๋ธ. 2์ดˆ ํด๋ฆฝ ๋‹จ์œ„ ๋ถ„๋ฅ˜์ด๊ณ  ์—ฃ์ง€ ๋””๋ฐ”์ด์Šค(๋ผ์ฆˆ๋ฒ ๋ฆฌํŒŒ์ด) ๋ฐฐํฌ๋ฅผ ์—ผ๋‘์— ๋’€๋‹ค.

  • ํด๋ž˜์Šค: baby_cry, bicycle, glass_break, gunshot
  • ๋ฐฑ๋ณธ: mn10_as (MobileNetV3, AudioSet ์‚ฌ์ „ํ•™์Šต) + MLP head
  • ์ž…๋ ฅ: 32,000 Hz mono, 2.0์ดˆ (64,000 ์ƒ˜ํ”Œ) โ†’ AugmentMelSTFT 128 mel
  • ํ•™์Šต ๋ฐ์ดํ„ฐ: retrina0678/miimo-audio-dataset
  • ์ฒดํฌํฌ์ธํŠธ ํฌ๊ธฐ: 17MB

์„ฑ๋Šฅ (held-out test, 5-fold)

metric mean std min max
accuracy 0.9774 0.0113 0.9617 0.9909
balanced accuracy 0.9763 0.0113 0.9618 0.9903
macro precision 0.9758 0.0123 0.9594 0.9906
macro recall 0.9763 0.0113 0.9618 0.9903
macro F1 0.9756 0.0122 0.9592 0.9904

ํด๋ž˜์Šค๋ณ„ recall (5-fold)

ํด๋ž˜์Šค mean std ๋น„๊ณ 
bicycle 0.9980 0.0046
baby_cry 0.9861 0.0113
gunshot 0.9849 0.0070
glass_break 0.9362 0.0314 โš ๏ธ ๊ฐ€์žฅ ์•ฝํ•จ โ€” ์ฃผ๋กœ gunshot ์œผ๋กœ ์˜ค๋ถ„๋ฅ˜

์ „์ฒด 5-fold pooled confusion matrix์—์„œ glass_break โ†’ gunshot ์˜ค๋ถ„๋ฅ˜๊ฐ€ 36๊ฑด์œผ๋กœ ๊ฐ€์žฅ ํฐ ์˜ค์ฐจ ์›์ธ์ด๋‹ค. ํŒŒ์—ด์Œ ๊ณ„์—ด์ด๋ผ 2์ดˆ ์ฐฝ์—์„œ ํ˜ผ๋™๋˜๋Š” ๊ฒƒ์œผ๋กœ ๋ณด์ธ๋‹ค.

์ถœ์ฒ˜๋ณ„ recall

ํด๋ž˜์Šค ์ถœ์ฒ˜ recall n
baby_cry AI Hub (์ฆ๊ฐ•) 0.9744 585
baby_cry ESC-50 (์ฆ๊ฐ•) 1.0000 305
baby_cry donateacry 1.0000 190
bicycle AI Hub (์ฆ๊ฐ•) 0.9980 490
glass_break AI Hub (์ฆ๊ฐ•) 0.9362 580
gunshot AI Hub (์ฆ๊ฐ•) 0.9849 595

ํŒŒ์ผ ๊ตฌ์„ฑ

best_model.pt        ์ตœ์ข… ์ฒดํฌํฌ์ธํŠธ (fold 2, best by macro_recall)
config.json          ํ•™์Šต ํ•˜์ดํผํŒŒ๋ผ๋ฏธํ„ฐ
labels.json          ํด๋ž˜์Šค ์ˆœ์„œ / label2id
all_metrics.json     fold๋ณ„ ์ „์ฒด ์ง€ํ‘œ
kfold_summary.csv    fold๋ณ„ ์š”์•ฝ
recall_by_source.csv ์ถœ์ฒ˜๋ณ„ recall
report.md            ์ƒ์„ธ ๋ฆฌํฌํŠธ (ํ˜ผ๋™ํ–‰๋ ฌ ํฌํ•จ)
fold_01..05/         fold๋ณ„ best.pt + ์˜ˆ์ธกยทํ˜ผ๋™ํ–‰๋ ฌ

best_model.pt ๋Š” dict์ด๊ณ  ํ‚ค๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™๋‹ค: model_state_dict, config, labels, label2id, id2label, fold, epoch, val_metrics

์‚ฌ์šฉ๋ฒ•

EfficientAT ์ €์žฅ์†Œ๊ฐ€ ํ•„์š”ํ•˜๋‹ค.

git clone https://github.com/fschmid56/EfficientAT
pip install torch torchaudio huggingface_hub
import sys, torch, torchaudio
from huggingface_hub import hf_hub_download

sys.path.insert(0, "EfficientAT")
from models.mn.model import get_model as get_mn
from models.preprocess import AugmentMelSTFT
from helpers.utils import NAME_TO_WIDTH

ckpt_path = hf_hub_download("retrina0678/miimo-efficientat-target4", "best_model.pt")
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
labels = ckpt["labels"]                      # ['baby_cry','bicycle','glass_break','gunshot']

model = get_mn(num_classes=len(labels), pretrained_name="mn10_as",
               width_mult=NAME_TO_WIDTH("mn10_as"), head_type="mlp")
model.load_state_dict(ckpt["model_state_dict"])
model.eval()

# ํ•™์Šต๊ณผ ๋™์ผํ•œ ์ „์ฒ˜๋ฆฌ โ€” freqm/timem ์€ ์ถ”๋ก  ์‹œ 0 ์œผ๋กœ ๋‘”๋‹ค
mel = AugmentMelSTFT(n_mels=128, sr=32000, win_length=800, hopsize=320, freqm=0, timem=0)
mel.eval()

wav, sr = torchaudio.load("clip.wav")
if sr != 32000:
    wav = torchaudio.functional.resample(wav, sr, 32000)
wav = wav.mean(0, keepdim=True)[:, :64000]   # mono, 2์ดˆ

with torch.no_grad():
    logits = model(mel(wav).unsqueeze(1))
    if isinstance(logits, (tuple, list)):
        logits = logits[0]
    probs = logits.flatten(1).softmax(-1)[0]

for lbl, p in sorted(zip(labels, probs.tolist()), key=lambda x: -x[1]):
    print(f"{lbl:<12} {p:.4f}")

ํ•™์Šต ์„ค์ •

ํ•ญ๋ชฉ ๊ฐ’
๋ฐฑ๋ณธ mn10_as (AudioSet ์‚ฌ์ „ํ•™์Šต)
head mlp
epochs / batch 20 / 16
optimizer lr 1e-4, weight decay 1e-4
freeze backbone ์•ž 2 epoch
CV StratifiedGroupKFold 5-fold, test_size 0.2
group column group_id (= aug_<source_file>)
selection metric macro_recall
class weight balanced + weighted sampler
seed 42
AMP on

๋ฐ์ดํ„ฐ ๋ˆ„์ˆ˜ ๋ฐฉ์ง€: ๊ฐ™์€ ์›๋ณธ ์Œ์›์—์„œ ๋‚˜์˜จ ํด๋ฆฝ์€ 50% ์˜ค๋ฒ„๋žฉ + ์ฆ๊ฐ•์œผ๋กœ ์„œ๋กœ ๊ฒน์น˜๋ฏ€๋กœ, ํด๋ฆฝ์ด ์•„๋‹ˆ๋ผ source_file ๋‹จ์œ„(group_id)๋กœ ๋ถ„ํ• ํ–ˆ๋‹ค.

ํด๋ž˜์Šค ๊ท ํ˜•

์ตœ์†Œ ํด๋ž˜์Šค(bicycle 3,747)์— ๋งž์ถฐ ํด๋ž˜์Šค๋‹น 3,747๊ฐœ๋กœ downsample โ†’ ์ด 14,988 ํด๋ฆฝ.

ํ•œ๊ณ„

  • 2์ดˆ ๊ณ ์ •์ฐฝ์ด๋ผ ๊ทธ๋ณด๋‹ค ๊ธด ์ด๋ฒคํŠธ์˜ ๋ฌธ๋งฅ์€ ๋ชป ๋ณธ๋‹ค.
  • glass_break recall์ด 0.936์œผ๋กœ ๊ฐ€์žฅ ๋‚ฎ๊ณ  fold ๊ฐ„ ํŽธ์ฐจ(std 0.031)๋„ ํฌ๋‹ค. gunshot ๊ณผ์˜ ํ˜ผ๋™์ด ์ฃผ ์›์ธ.
  • ํ•™์Šต ๋ฐ์ดํ„ฐ๊ฐ€ AI Hub ์ค‘์‹ฌ์ด๋ผ ๋…น์Œ ํ™˜๊ฒฝ์ด ๋‹ค๋ฅธ ์‹ค์‚ฌ์šฉ ํ™˜๊ฒฝ์—์„œ๋Š” ์„ฑ๋Šฅ์ด ๋–จ์–ด์งˆ ์ˆ˜ ์žˆ๋‹ค. recall_by_source.csv ์ฐธ๊ณ .
  • 4๊ฐœ ํด๋ž˜์Šค๋งŒ ๋‹ค๋ฃจ๊ณ , ๊ทธ ์™ธ ์†Œ๋ฆฌ๋Š” ๊ฐ€์žฅ ๊ฐ€๊นŒ์šด ํด๋ž˜์Šค๋กœ ๊ฐ•์ œ ๋ถ„๋ฅ˜๋œ๋‹ค. ์‹ค์‚ฌ์šฉ ์‹œ ํ™•๋ฅ  ์ž„๊ณ„๊ฐ’(threshold)์„ ๋‘๊ณ  reject ์ฒ˜๋ฆฌํ•˜๋Š” ๊ฒƒ์„ ๊ถŒํ•œ๋‹ค.

๋ผ์ด์„ ์Šค

โš ๏ธ ์ƒ์—…์  ์ด์šฉ ๋ถˆ๊ฐ€.

ํ•™์Šต ๋ฐ์ดํ„ฐ์— ESC-50(CC BY-NC 3.0)๊ณผ AI Hub ์ œ๊ณต ๋ฐ์ดํ„ฐ(์žฌ๋ฐฐํฌ ์ œํ•œ)๊ฐ€ ํฌํ•จ๋˜์–ด ์žˆ๋‹ค. ์ด ๊ฐ€์ค‘์น˜๋Š” ํ•ด๋‹น ๋ฐ์ดํ„ฐ์—์„œ ํŒŒ์ƒ๋˜์—ˆ์œผ๋ฏ€๋กœ ์› ๋ฐ์ดํ„ฐ์˜ ์ œ์•ฝ์„ ๊ทธ๋Œ€๋กœ ์Šน๊ณ„ํ•œ๋‹ค.

์ธ์šฉ

EfficientAT: F. Schmid, K. Koutini, G. Widmer.
"Efficient Large-Scale Audio Tagging via Transformer-to-CNN Knowledge Distillation." ICASSP 2023.

ESC-50: K. J. Piczak. "ESC: Dataset for Environmental Sound Classification." ACM MM 2015.