| --- |
| language: |
| - en |
| - zh |
| - id |
| - ms |
| base_model: openai/whisper-large-v3 |
| tags: |
| - language-identification |
| - code-switching |
| - audio-classification |
| - whisper |
| pipeline_tag: audio-classification |
| --- |
| |
| # CS-LID-4 |
|
|
| Utterance-level **code-switch detection** and **4-language identification** |
| ({English, Mandarin, Indonesian, Malay}) for conversational speech. |
|
|
| The model is a pair of linear heads on pooled |
| [openai/whisper-large-v3](https://huggingface.co/openai/whisper-large-v3) |
| encoder features: |
|
|
| - `cs` head — 1 logit (sigmoid): probability the utterance contains |
| code-switching. |
| - `lang4` head — 4-way softmax over `{en, zh, id, ms}`: the (dominant) |
| language of the utterance. Trained on monolingual utterances. |
|
|
| Each head comes in two pooling variants stored in the same checkpoint: |
|
|
| - `real` (recommended) — mean over the encoder frames that correspond to |
| actual audio (50 frames/s, up to 1500). |
| - `pad` — mean over all 1500 encoder frames of the padded 30 s window. |
|
|
| Training data spans conversational and read speech in the four languages |
| plus en–zh, en–id, and en–ms code-switching corpora |
| (see `yyhenggg/CS-LID-Dataset`), with utterance-level labels. |
|
|
| ## Usage |
|
|
| ```bash |
| pip install torch transformers soundfile huggingface_hub |
| python inference.py audio1.wav audio2.flac |
| ``` |
|
|
| Audio should be mono; 16 kHz is expected (other rates are resampled when |
| `torchaudio` is installed) and input longer than 30 s is truncated. |
| Output, one line per file: |
|
|
| ``` |
| audio1.wav P(cs)=0.912 lang=zh p_en=0.031 p_zh=0.952 p_id=0.009 p_ms=0.008 |
| ``` |
|
|
| ### Python API |
|
|
| ```python |
| import torch |
| from huggingface_hub import hf_hub_download |
| from transformers import WhisperProcessor, WhisperForConditionalGeneration |
| |
| ck = torch.load(hf_hub_download('yyhenggg/CS-LID-Model', 'cslid4_head.pt'), |
| map_location='cpu', weights_only=True) |
| sd, langs = ck['head_state_dict'], ck['langs'] # langs = ['en','zh','id','ms'] |
| |
| proc = WhisperProcessor.from_pretrained('openai/whisper-large-v3') |
| enc = WhisperForConditionalGeneration.from_pretrained( |
| 'openai/whisper-large-v3').eval().model.encoder |
| |
| audio = ... # float32 mono 16 kHz numpy array, <= 30 s |
| feats = proc([audio], sampling_rate=16000, return_tensors='pt', |
| padding='max_length').input_features |
| with torch.no_grad(): |
| h = enc(feats, return_dict=True).last_hidden_state.float()[0] |
| n_real = min(int(len(audio) / 16000 * 50), 1500) # 'real' pooling |
| x = h[:max(n_real, 1)].mean(0) |
| |
| p_cs = torch.sigmoid(x @ sd['cs_real.weight'].T + sd['cs_real.bias']).item() |
| p_lang = torch.softmax(x @ sd['lang4_real.weight'].T |
| + sd['lang4_real.bias'], dim=-1) |
| print(p_cs, dict(zip(langs, p_lang.tolist()))) |
| ``` |
|
|
| ## Checkpoint format |
|
|
| `cslid4_head.pt` is a plain `torch.save` dict: |
|
|
| ``` |
| head_state_dict: |
| cs_real.weight [1, 1280] cs_real.bias [1] |
| cs_pad.weight [1, 1280] cs_pad.bias [1] |
| lang4_real.weight [4, 1280] lang4_real.bias [4] |
| lang4_pad.weight [4, 1280] lang4_pad.bias [4] |
| langs: ['en', 'zh', 'id', 'ms'] |
| design: short provenance string |
| ``` |
|
|