retrina0678's picture
Add model card / clean paths
db48d32 verified
|
Raw
History Blame Contribute Delete
7.04 kB
---
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](https://github.com/fschmid56/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](https://huggingface.co/datasets/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 ์ €์žฅ์†Œ๊ฐ€ ํ•„์š”ํ•˜๋‹ค.
```bash
git clone https://github.com/fschmid56/EfficientAT
pip install torch torchaudio huggingface_hub
```
```python
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](https://github.com/karolpiczak/ESC-50)(**CC BY-NC 3.0**)๊ณผ
AI Hub ์ œ๊ณต ๋ฐ์ดํ„ฐ(**์žฌ๋ฐฐํฌ ์ œํ•œ**)๊ฐ€ ํฌํ•จ๋˜์–ด ์žˆ๋‹ค. ์ด ๊ฐ€์ค‘์น˜๋Š” ํ•ด๋‹น ๋ฐ์ดํ„ฐ์—์„œ
ํŒŒ์ƒ๋˜์—ˆ์œผ๋ฏ€๋กœ ์› ๋ฐ์ดํ„ฐ์˜ ์ œ์•ฝ์„ ๊ทธ๋Œ€๋กœ ์Šน๊ณ„ํ•œ๋‹ค.
- ์ƒ์—…์  ์šฉ๋„๋กœ๋Š” ์‚ฌ์šฉํ•  ์ˆ˜ ์—†๋‹ค.
- ๋ฐฑ๋ณธ `mn10_as` ์ž์ฒด์˜ ๋ผ์ด์„ ์Šค๋Š” [EfficientAT ์ €์žฅ์†Œ](https://github.com/fschmid56/EfficientAT)๋ฅผ ๋”ฐ๋ฅธ๋‹ค.
- ์ž์„ธํ•œ ์ถœ์ฒ˜๋Š” [๋ฐ์ดํ„ฐ์…‹ ์นด๋“œ](https://huggingface.co/datasets/retrina0678/miimo-audio-dataset)๋ฅผ ์ฐธ๊ณ .
## ์ธ์šฉ
```
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.
```