| --- |
| 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. |
| ``` |
| |