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_breakrecall์ด 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 ์ ๊ณต ๋ฐ์ดํฐ(์ฌ๋ฐฐํฌ ์ ํ)๊ฐ ํฌํจ๋์ด ์๋ค. ์ด ๊ฐ์ค์น๋ ํด๋น ๋ฐ์ดํฐ์์ ํ์๋์์ผ๋ฏ๋ก ์ ๋ฐ์ดํฐ์ ์ ์ฝ์ ๊ทธ๋๋ก ์น๊ณํ๋ค.
- ์์ ์ ์ฉ๋๋ก๋ ์ฌ์ฉํ ์ ์๋ค.
- ๋ฐฑ๋ณธ
mn10_as์์ฒด์ ๋ผ์ด์ ์ค๋ EfficientAT ์ ์ฅ์๋ฅผ ๋ฐ๋ฅธ๋ค. - ์์ธํ ์ถ์ฒ๋ ๋ฐ์ดํฐ์ ์นด๋๋ฅผ ์ฐธ๊ณ .
์ธ์ฉ
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.