frankie137 commited on
Commit
b94bf59
·
verified ·
1 Parent(s): 4ca6408

Add files using upload-large-folder tool

Browse files
Files changed (4) hide show
  1. README.md +96 -3
  2. config.json +4 -4
  3. configuration_alm2vec.py +5 -0
  4. modeling_alm2vec.py +1118 -0
README.md CHANGED
@@ -1,6 +1,8 @@
1
  ---
 
2
  license: apache-2.0
3
  language:
 
4
  - en
5
  pipeline_tag: feature-extraction
6
  tags:
@@ -9,11 +11,28 @@ tags:
9
  - custom_code
10
  - multimodal
11
  base_model: mispeech/midashenglm-7b-0804-fp32
 
12
  ---
13
 
14
- # AudioEmb (Finetune)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
 
16
- Audio-text embedding model for retrieval, based on [MiDashengLM](https://huggingface.co/mispeech/midashenglm-7b-0804-fp32). Checkpoint: **finetune** stage.
17
 
18
  Requirements: `transformers>=4.52`, `torch`, `safetensors`, and `torchaudio` for non-WAV audio. Requires a GPU (~31GB weights) and `trust_remote_code=True`.
19
 
@@ -23,7 +42,7 @@ Requirements: `transformers>=4.52`, `torch`, `safetensors`, and `torchaudio` for
23
  import torch
24
  from transformers import AutoModel, AutoTokenizer
25
 
26
- repo_id = "cara-ai/AudioEmb-finetune"
27
  tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)
28
  model = AutoModel.from_pretrained(
29
  repo_id, trust_remote_code=True, torch_dtype=torch.float32
@@ -41,3 +60,77 @@ similarity = query_emb @ doc_emb.T
41
  print(similarity)
42
  ```
43
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+
3
  license: apache-2.0
4
  language:
5
+
6
  - en
7
  pipeline_tag: feature-extraction
8
  tags:
 
11
  - custom_code
12
  - multimodal
13
  base_model: mispeech/midashenglm-7b-0804-fp32
14
+
15
  ---
16
 
17
+ <h1 align="center">ALM2Vec-FT</h1>
18
+
19
+ <p align="center">
20
+ <a href="https://arxiv.org/abs/xxxx.xxxxx"><img src="https://img.shields.io/badge/Paper-arXiv-b31b1b?logo=arxiv&logoColor=white" alt="Paper"></a>
21
+ <a href="https://caml-labs.github.io/ALM2Vec"><img src="https://img.shields.io/badge/Project-Page-1f6feb?logo=googlechrome&logoColor=white" alt="Project Page"></a>
22
+ <a href="https://github.com/caml-labs/ALM2Vec"><img src="https://img.shields.io/badge/Code-GitHub-181717?logo=github&logoColor=white" alt="GitHub"></a>
23
+ </p>
24
+
25
+ **ALM2Vec** is a universal audio embedding model for retrieval, derived from a pretrained large audio–language model (LALM). Instead of being optimized only for audio–caption matching like conventional contrastive dual-encoders, it transfers the audio understanding, instruction-following, and reasoning abilities of LALMs into a single unified embedding space that works across audio domains, task types, and user intents.
26
+
27
+ Its key feature is **instruction-aware retrieval**: a natural-language instruction guides the embedding, so the *same* audio can be encoded differently for different needs. This supports:
28
+
29
+ - **Instruction-aware retrieval** — focus the embedding on a specific aspect of the audio.
30
+ - **Text ↔ audio retrieval** — bidirectional matching between audio and text.
31
+ - **Audio question answering** — match an audio query plus a question against candidate answers.
32
+
33
+ ALM2Vec achieves competitive results on standard audio and speech retrieval benchmarks while adding these controllable retrieval capabilities. See the [project page](https://caml-labs.github.io/ALM2Vec/) for interactive demos.
34
 
35
+ This repository hosts the **finetune** checkpoint, built on [MiDashengLM](https://huggingface.co/mispeech/midashenglm-7b-0804-fp32).
36
 
37
  Requirements: `transformers>=4.52`, `torch`, `safetensors`, and `torchaudio` for non-WAV audio. Requires a GPU (~31GB weights) and `trust_remote_code=True`.
38
 
 
42
  import torch
43
  from transformers import AutoModel, AutoTokenizer
44
 
45
+ repo_id = "cara-ai/ALM2Vec-FT"
46
  tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True)
47
  model = AutoModel.from_pretrained(
48
  repo_id, trust_remote_code=True, torch_dtype=torch.float32
 
60
  print(similarity)
61
  ```
62
 
63
+ ## Results
64
+
65
+ **ALM2Vec-FT** is the checkpoint hosted in this repository; **ALM2Vec-PT** is the pretrain variant. In every table, **bold** marks the best score and <u>underline</u> the second best.
66
+
67
+ ### Text–audio retrieval — AudioCaps
68
+
69
+ | Method | T→A R@1 | T→A R@5 | T→A R@10 | A→T R@1 | A→T R@5 | A→T R@10 |
70
+ | --- | --- | --- | --- | --- | --- | --- |
71
+ | LAION-CLAP | 36.1 | 71.8 | 83.9 | 46.8 | <u>82.9</u> | <u>90.7</u> |
72
+ | MS-CLAP | 15.4 | 47.2 | 64.5 | 32.0 | 66.0 | 79.2 |
73
+ | WavCaps-CLAP-PT | 39.7 | 74.5 | 86.1 | 51.7 | 82.3 | 90.6 |
74
+ | WavCaps-CLAP-FT | <u>42.2</u> | <u>76.5</u> | <u>87.1</u> | <u>54.6</u> | **85.2** | **92.4** |
75
+ | JINA-Embed.-v5 | 20.4 | 50.3 | 64.4 | 23.1 | 52.7 | 67.2 |
76
+ | **ALM2Vec-PT** | 40.0 | 74.5 | 85.9 | 43.8 | 74.3 | 86.5 |
77
+ | **ALM2Vec-FT** | **43.2** | **78.0** | **87.8** | **55.5** | 80.0 | 88.2 |
78
+
79
+ ### Text–audio retrieval — Clotho
80
+
81
+ | Method | T→A R@1 | T→A R@5 | T→A R@10 | A→T R@1 | A→T R@5 | A→T R@10 |
82
+ | --- | --- | --- | --- | --- | --- | --- |
83
+ | LAION-CLAP | 16.1 | 38.3 | 51.1 | 22.7 | 48.5 | 60.8 |
84
+ | MS-CLAP | 15.6 | 38.9 | 51.4 | 22.1 | 48.9 | 62.0 |
85
+ | WavCaps-CLAP-PT | 19.5 | 45.2 | 58.2 | 23.4 | 50.9 | 63.4 |
86
+ | WavCaps-CLAP-FT | <u>19.7</u> | <u>45.7</u> | <u>59.4</u> | <u>26.9</u> | <u>52.6</u> | <u>64.9</u> |
87
+ | JINA-Embed.-v5 | 9.2 | 23.9 | 35.0 | 10.5 | 24.7 | 34.3 |
88
+ | **ALM2Vec-PT** | 19.2 | 43.4 | 55.7 | 17.9 | 39.4 | 52.2 |
89
+ | **ALM2Vec-FT** | **24.8** | **52.9** | **65.8** | **27.9** | **52.7** | **66.3** |
90
+
91
+
92
+ ### Speech retrieval — LibriSQA
93
+
94
+
95
+ | Method | T→S R@1 | T→S R@5 | T→S R@10 | S→T R@1 | S→T R@5 | S→T R@10 |
96
+ | -------------- | --------------- | -------- | -------- | --------------- | -------- | -------- |
97
+ | LAION-CLAP † | 0.0 | 0.1 | 0.8 | 0.1 | 0.2 | 0.6 |
98
+ | Whisper+BGE | 83.7 | 93.3 | 94.9 | 85.2 | 93.4 | 95.3 |
99
+ | CLSR | **85.0** | <u>93.4</u> | <u>95.0</u> | <u>85.5</u> | <u>94.0</u> | <u>95.6</u> |
100
+ | **ALM2Vec-PT** | 43.7 | 64.5 | 72.8 | 11.2 | 24.9 | 34.1 |
101
+ | **ALM2Vec-FT** | <u>84.7</u> | **94.1** | **95.8** | **86.0** | **95.2** | **97.2** |
102
+
103
+
104
+ ### Audio understanding — MMAU-mini (accuracy)
105
+
106
+
107
+ | Method | Overall | Music | Sound | Speech |
108
+ | ------------------ | -------- | -------- | -------- | -------- |
109
+ | GPT-4o Audio ‡ | 60.8 | 63.2 | 64.6 | 56.3 |
110
+ | Gemini 2.5 Pro ‡ | <u>71.6</u> | <u>75.1</u> | 71.5 | 68.3 |
111
+ | Qwen2.5-Omni ‡ | 71.5 | 65.9 | <u>78.1</u> | <u>70.6</u> |
112
+ | Audio Flamingo 3 ‡ | **73.1** | **76.9** | 66.1 | **73.9** |
113
+ | **ALM2Vec-PT** | 66.3 | 62.3 | **78.7** | 58.0 |
114
+ | **ALM2Vec-FT** | 63.0 | 61.7 | 74.8 | 52.6 |
115
+
116
+
117
+ † LAION-CLAP is not trained for speech and effectively fails on LibriSQA; shown for reference.
118
+ ‡ Generative large audio–language models, listed as reference upper bounds rather than directly comparable retrieval baselines.
119
+
120
+ ## Citation
121
+
122
+ If you find this work useful, please consider citing:
123
+
124
+ ```
125
+ @article{ALM2Vec2026,
126
+ title={ALM2Vec: Learning Audio Embeddings for Universal
127
+ Audio Retrieval with Large Audio-Language Models},
128
+ author={TBD},
129
+ journal={arXiv preprint arXiv:TBD},
130
+ year={2026}
131
+ }
132
+ ```
133
+
134
+ ## Acknowledgement
135
+
136
+ ALM2Vec is built on [MiDashengLM](https://github.com/xiaomi-research/dasheng-lm) and further trained for universal audio retrieval. We thank MiDashengLM and its underlying [Dasheng](https://github.com/RicherMans/Dasheng) audio encoder for their open-source contributions.
config.json CHANGED
@@ -1,12 +1,12 @@
1
  {
2
- "model_type": "audio_emb",
3
  "architectures": [
4
- "AudioEmbModel"
5
  ],
6
  "torch_dtype": "float32",
7
  "transformers_version": "5.0.0.dev0",
8
  "auto_map": {
9
- "AutoConfig": "configuration_audio_emb.AudioEmbConfig",
10
- "AutoModel": "modeling_audio_emb.AudioEmbModel"
11
  }
12
  }
 
1
  {
2
+ "model_type": "alm2vec",
3
  "architectures": [
4
+ "ALM2VecModel"
5
  ],
6
  "torch_dtype": "float32",
7
  "transformers_version": "5.0.0.dev0",
8
  "auto_map": {
9
+ "AutoConfig": "configuration_alm2vec.ALM2VecConfig",
10
+ "AutoModel": "modeling_alm2vec.ALM2VecModel"
11
  }
12
  }
configuration_alm2vec.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ from transformers import PreTrainedConfig
2
+
3
+
4
+ class ALM2VecConfig(PreTrainedConfig):
5
+ model_type = "alm2vec"
modeling_alm2vec.py ADDED
@@ -0,0 +1,1118 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ALM2Vec audio-text embedding model."""
2
+
3
+ import math
4
+ from torch.utils.checkpoint import checkpoint
5
+ import wave
6
+ from io import BytesIO
7
+ from pathlib import Path
8
+ from tempfile import NamedTemporaryFile
9
+ from urllib.parse import urlparse
10
+ from urllib.request import urlopen
11
+
12
+ from .configuration_alm2vec import ALM2VecConfig
13
+
14
+ import collections
15
+ import collections.abc
16
+
17
+ from dataclasses import dataclass
18
+ import torch
19
+ import torch.nn as nn
20
+ import torchaudio.functional as F
21
+ from torch import Tensor
22
+ from torch.nn.functional import scaled_dot_product_attention
23
+ from typing import Any, Dict, Callable, Iterable, List, Optional, Sequence, Tuple, Union, cast
24
+
25
+ from transformers import PreTrainedModel, PreTrainedConfig, GenerationMixin
26
+ from transformers import AutoTokenizer
27
+
28
+ from transformers.models.qwen2_5_omni.configuration_qwen2_5_omni import (
29
+ Qwen2_5OmniTextConfig,
30
+ )
31
+ from transformers.models.qwen2_5_omni.modeling_qwen2_5_omni import (
32
+ Qwen2_5OmniThinkerTextModel,
33
+ )
34
+ from transformers.cache_utils import Cache
35
+ from transformers.modeling_outputs import BaseModelOutputWithPast, ModelOutput
36
+ from transformers.utils import can_return_tuple
37
+
38
+ import copy
39
+
40
+ try:
41
+ import torchaudio
42
+ except ImportError:
43
+ torchaudio = None
44
+
45
+ _Tuple2 = Union[int, Tuple[int, int], Sequence[int]]
46
+
47
+ TARGET_SR = 16000
48
+
49
+ QUERY_INSTRUCTION = "Based on the question asked in the text query and context in the audio query, retrieve the relevant text document associated with that question."
50
+ DOC_INSTRUCTION = "Represent the user's input."
51
+
52
+
53
+ def _resolve_tuple2(x: _Tuple2) -> Tuple[int, int]:
54
+ if isinstance(x, collections.abc.Sequence):
55
+ assert len(x) == 2, (
56
+ f"Expected a sequence of length 2, got {x} with length {len(x)}"
57
+ )
58
+ return cast(Tuple[int, int], tuple(x))
59
+ return (x, x)
60
+
61
+
62
+
63
+
64
+ DASHENG_ARCH_CONFIG = {
65
+ "audio_encoder_config": {
66
+ "attn_drop_rate": 0.0,
67
+ "center": True,
68
+ "depth": 32,
69
+ "drop_rate": 0.0,
70
+ "embed_dim": 1280,
71
+ "f_max": 8000.0,
72
+ "f_min": 0.0,
73
+ "hop_length": 160,
74
+ "init_values": None,
75
+ "input_channels": 1,
76
+ "mlp_ratio": 4.0,
77
+ "model_type": "midashenglm_dasheng_encoder",
78
+ "n_fft": 512,
79
+ "n_mels": 64,
80
+ "num_heads": 16,
81
+ "outputdim": 527,
82
+ "patch_size": [
83
+ 64,
84
+ 4
85
+ ],
86
+ "patch_stride": [
87
+ 64,
88
+ 4
89
+ ],
90
+ "qkv_bias": True,
91
+ "sample_rate": 16000,
92
+ "target_length": 1008,
93
+ "win_length": 512
94
+ },
95
+
96
+ "audio_projector_config": {
97
+ "in_dim": 1280,
98
+ "downsample_rate": 5,
99
+ "out_dim": 3584,
100
+ },
101
+
102
+ "text_config": {
103
+ "attention_dropout": 0.0,
104
+ "hidden_act": "silu",
105
+ "hidden_size": 3584,
106
+ "init_std": 0.02,
107
+ "initializer_range": 0.02,
108
+ "intermediate_size": 18944,
109
+ "max_position_embeddings": 32768,
110
+ "max_window_layers": 28,
111
+ "model_type": "qwen2_5_omni_text",
112
+ "num_attention_heads": 28,
113
+ "num_hidden_layers": 28,
114
+ "num_key_value_heads": 4,
115
+ "rms_norm_eps": 1e-06,
116
+ "rope_scaling": {
117
+ "mrope_section": [
118
+ 16,
119
+ 24,
120
+ 24
121
+ ],
122
+ "rope_type": "default",
123
+ "type": "default"
124
+ },
125
+ "rope_theta": 1000000.0,
126
+ "sliding_window": 32768,
127
+ "use_cache": True,
128
+ "use_sliding_window": False,
129
+ "vocab_size": 152064
130
+ },
131
+
132
+ "lite_random_decoder_config": {
133
+ "attention_dropout": 0.0,
134
+ "hidden_act": "silu",
135
+ "hidden_size": 576,
136
+ "init_std": 0.02,
137
+ "initializer_range": 0.02,
138
+ "intermediate_size": 1536,
139
+ "max_position_embeddings": 2048,
140
+ "max_window_layers": 12,
141
+ "model_type": "qwen2_5_omni_text",
142
+ "num_attention_heads": 8,
143
+ "num_hidden_layers": 12,
144
+ "num_key_value_heads": 4,
145
+ "rms_norm_eps": 1e-06,
146
+ "rope_scaling": {
147
+ "mrope_section": [
148
+ 12,
149
+ 12,
150
+ 12
151
+ ],
152
+ "rope_type": "default",
153
+ "type": "default"
154
+ },
155
+ "rope_theta": 1000000.0,
156
+ "sliding_window": 2048,
157
+ "use_cache": True,
158
+ "use_sliding_window": False,
159
+ "vocab_size": 152064
160
+ }
161
+ }
162
+
163
+
164
+ class DashengConfig(PreTrainedConfig):
165
+ model_type = "midashenglm_dasheng_encoder"
166
+
167
+ def __init__(
168
+ self,
169
+ embed_dim: int = 768,
170
+ outputdim: int = 527,
171
+ patch_size: Union[int, Tuple[int, int]] = 16,
172
+ patch_stride: Union[int, Tuple[int, int]] = 16,
173
+ input_channels: int = 1,
174
+ target_length: int = 1012,
175
+ depth: int = 12,
176
+ num_heads: int = 12,
177
+ mlp_ratio: float = 4.0,
178
+ qkv_bias: bool = True,
179
+ init_values: Optional[float] = None,
180
+ drop_rate: float = 0.0,
181
+ attn_drop_rate: float = 0.0,
182
+ f_min: float = 0.0,
183
+ f_max: float = 8000.0,
184
+ center: bool = True,
185
+ win_length: int = 512,
186
+ hop_length: int = 160,
187
+ sample_rate: int = 16000,
188
+ n_fft: int = 512,
189
+ n_mels: int = 64,
190
+ **kwargs,
191
+ ):
192
+ self.embed_dim = embed_dim
193
+ self.outputdim = outputdim
194
+ self.patch_size = patch_size
195
+ self.patch_stride = patch_stride
196
+ self.input_channels = input_channels
197
+ self.target_length = target_length
198
+ self.depth = depth
199
+ self.num_heads = num_heads
200
+ self.mlp_ratio = mlp_ratio
201
+ self.qkv_bias = qkv_bias
202
+ self.init_values = init_values
203
+ self.drop_rate = drop_rate
204
+ self.attn_drop_rate = attn_drop_rate
205
+ self.f_min = f_min
206
+ self.f_max = f_max
207
+ self.center = center
208
+ self.win_length = win_length
209
+ self.hop_length = hop_length
210
+ self.sample_rate = sample_rate
211
+ self.n_fft = n_fft
212
+ self.n_mels = n_mels
213
+ super().__init__(**kwargs)
214
+
215
+
216
+ class AudioPatchEmbed(nn.Module):
217
+ def __init__(
218
+ self,
219
+ input_size: _Tuple2 = 64,
220
+ patch_size: _Tuple2 = 16,
221
+ patch_stride: _Tuple2 = 16,
222
+ in_chans: int = 1,
223
+ embed_dim: int = 768,
224
+ norm_layer: Optional[Callable] = None,
225
+ flatten: bool = False,
226
+ ):
227
+ super().__init__()
228
+ self.input_size = _resolve_tuple2(input_size)
229
+ self.patch_size = _resolve_tuple2(patch_size)
230
+ self.patch_stride = _resolve_tuple2(patch_stride)
231
+ self.grid_size = (
232
+ self.input_size[0] // self.patch_stride[0],
233
+ self.input_size[1] // self.patch_stride[1],
234
+ )
235
+ self.num_patches = self.grid_size[0] * self.grid_size[1]
236
+ self.flatten = flatten
237
+
238
+ self.proj = nn.Conv2d(
239
+ in_chans,
240
+ embed_dim,
241
+ kernel_size=self.patch_size,
242
+ stride=self.patch_stride,
243
+ )
244
+ self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
245
+
246
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
247
+ x = self.proj(x)
248
+ if self.flatten:
249
+ x = torch.permute(
250
+ torch.flatten(x, 2, 3), (0, 2, 1)
251
+ ) # rearrange(x, "b c f t -> b (f t) c")
252
+ x = self.norm(x)
253
+ return x
254
+
255
+
256
+ class LayerScale(nn.Module):
257
+ def __init__(self, dim, init_values=1e-5, inplace=False):
258
+ super().__init__()
259
+ self.inplace = inplace
260
+ self.gamma = nn.Parameter(init_values * torch.ones(dim))
261
+
262
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
263
+ return x.mul_(self.gamma) if self.inplace else x * self.gamma
264
+
265
+
266
+ class DashengMlp(nn.Module):
267
+ def __init__(
268
+ self,
269
+ in_features: int,
270
+ hidden_features: Optional[int] = None,
271
+ out_features: Optional[int] = None,
272
+ drop: float = 0.0,
273
+ ):
274
+ super().__init__()
275
+ out_features = out_features or in_features
276
+ hidden_features = hidden_features or in_features
277
+ self.fc1 = nn.Linear(in_features, hidden_features)
278
+ self.act = nn.GELU()
279
+ self.fc2 = nn.Linear(hidden_features, out_features)
280
+ self.drop = nn.Dropout(drop)
281
+
282
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
283
+ x = self.fc1(x)
284
+ x = self.act(x)
285
+ x = self.drop(x)
286
+ x = self.fc2(x)
287
+ x = self.drop(x)
288
+ return x
289
+
290
+
291
+ class DashengAttention(nn.Module):
292
+ def __init__(
293
+ self,
294
+ dim: int,
295
+ num_heads: int = 8,
296
+ qkv_bias: bool = False,
297
+ attn_drop: float = 0.0,
298
+ proj_drop: float = 0.0,
299
+ ):
300
+ super().__init__()
301
+ assert dim % num_heads == 0, "dim should be divisible by num_heads"
302
+ self.num_heads = num_heads
303
+ head_dim = dim // num_heads
304
+ self.scale = head_dim**-0.5
305
+
306
+ self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
307
+ self.attn_drop = nn.Dropout(attn_drop)
308
+ self.proj = nn.Linear(dim, dim)
309
+ self.proj_drop = nn.Dropout(proj_drop)
310
+
311
+ def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None):
312
+ B, N, C = x.shape
313
+ q, k, v = (
314
+ self.qkv(x)
315
+ .reshape(B, N, 3, self.num_heads, C // self.num_heads)
316
+ .permute(2, 0, 3, 1, 4)
317
+ .unbind(0)
318
+ )
319
+ x = scaled_dot_product_attention(
320
+ q,
321
+ k,
322
+ v,
323
+ attn_mask=mask[:, None, None, :] if mask is not None else None,
324
+ )
325
+ x = x.transpose(1, 2).reshape(B, N, C)
326
+ x = self.proj(x)
327
+ x = self.proj_drop(x)
328
+ return x
329
+
330
+
331
+ class DashengBlock(nn.Module):
332
+ def __init__(
333
+ self,
334
+ dim: int,
335
+ num_heads: int,
336
+ mlp_ratio: float = 4.0,
337
+ qkv_bias: bool = False,
338
+ drop: float = 0.0,
339
+ attn_drop: float = 0.0,
340
+ init_values: Optional[float] = None,
341
+ ):
342
+ super().__init__()
343
+ self.norm1 = nn.LayerNorm(dim, eps=1e-6)
344
+ self.attn = DashengAttention(
345
+ dim,
346
+ num_heads=num_heads,
347
+ qkv_bias=qkv_bias,
348
+ attn_drop=attn_drop,
349
+ proj_drop=drop,
350
+ )
351
+ self.ls1 = (
352
+ LayerScale(dim, init_values=init_values) if init_values else nn.Identity()
353
+ )
354
+
355
+ self.norm2 = nn.LayerNorm(dim, eps=1e-6)
356
+ self.mlp = DashengMlp(
357
+ in_features=dim,
358
+ hidden_features=int(dim * mlp_ratio),
359
+ drop=drop,
360
+ )
361
+ self.ls2 = (
362
+ LayerScale(dim, init_values=init_values) if init_values else nn.Identity()
363
+ )
364
+
365
+ # Kwargs usually has a mask parameter that is passed to Attention
366
+ def forward(
367
+ self,
368
+ x: torch.Tensor,
369
+ mask: Optional[torch.Tensor] = None,
370
+ ) -> torch.Tensor:
371
+ x = x + self.ls1(self.attn(self.norm1(x), mask))
372
+ x = x + self.ls2(self.mlp(self.norm2(x)))
373
+ return x
374
+
375
+
376
+ class DashengFrontend(nn.Module):
377
+ def __init__(self, config: DashengConfig):
378
+ super().__init__()
379
+ self.config = config
380
+
381
+ spectrogram_window, melscale_fbanks = self._build_frontend_buffers()
382
+ self.register_buffer(
383
+ "spectrogram_window",
384
+ spectrogram_window,
385
+ persistent=False,
386
+ )
387
+ self.spectrogram_window: torch.Tensor
388
+ self.register_buffer("melscale_fbanks", melscale_fbanks, persistent=False)
389
+ self.melscale_fbanks: torch.Tensor
390
+
391
+ def _build_frontend_buffers(self) -> tuple[torch.Tensor, torch.Tensor]:
392
+ # Build on CPU explicitly: from_pretrained may construct modules on meta device.
393
+ with torch.device("cpu"):
394
+ spectrogram_window = torch.hann_window(
395
+ self.config.win_length,
396
+ dtype=torch.float32,
397
+ )
398
+ melscale_fbanks = F.melscale_fbanks(
399
+ n_freqs=self.config.n_fft // 2 + 1,
400
+ f_min=self.config.f_min,
401
+ f_max=self.config.f_max,
402
+ n_mels=self.config.n_mels,
403
+ sample_rate=self.config.sample_rate,
404
+ ).to(torch.float32)
405
+ return spectrogram_window, melscale_fbanks
406
+
407
+ def ensure_frontend_buffers(self, device: torch.device) -> None:
408
+ """Self-heal non-persistent audio frontend buffers if corrupted/uninitialized."""
409
+ expected_win_shape = (self.config.win_length,)
410
+ expected_fb_shape = (self.config.n_fft // 2 + 1, self.config.n_mels)
411
+
412
+ def _is_bad(name: str, tensor: torch.Tensor, expected_shape: tuple[int, ...]) -> bool:
413
+ if tensor is None:
414
+ return True
415
+ if getattr(tensor, "is_meta", False):
416
+ return True
417
+ if tuple(tensor.shape) != expected_shape:
418
+ return True
419
+ t = tensor.detach().float()
420
+ if not torch.isfinite(t).all().item():
421
+ return True
422
+ if t.numel() > 0 and t.abs().max().item() > 1e6:
423
+ return True
424
+ return False
425
+
426
+ win_bad = _is_bad("spectrogram_window", self.spectrogram_window, expected_win_shape)
427
+ fb_bad = _is_bad("melscale_fbanks", self.melscale_fbanks, expected_fb_shape)
428
+ if win_bad or fb_bad:
429
+ new_win, new_fb = self._build_frontend_buffers()
430
+ self.spectrogram_window = new_win.to(device=device)
431
+ self.melscale_fbanks = new_fb.to(device=device)
432
+ print(
433
+ f"[WARN] Rebuilt frontend buffers (win_bad={win_bad}, fb_bad={fb_bad})",
434
+ flush=True,
435
+ )
436
+ else:
437
+ if self.spectrogram_window.device != device:
438
+ self.spectrogram_window = self.spectrogram_window.to(device=device)
439
+ if self.melscale_fbanks.device != device:
440
+ self.melscale_fbanks = self.melscale_fbanks.to(device=device)
441
+
442
+ def forward(self, waveform: torch.Tensor) -> torch.Tensor:
443
+ self.ensure_frontend_buffers(waveform.device)
444
+
445
+ spectrogram = F.spectrogram(
446
+ waveform=waveform.to(torch.float32),
447
+ pad=0,
448
+ window=self.spectrogram_window,
449
+ n_fft=self.config.n_fft,
450
+ hop_length=self.config.hop_length,
451
+ win_length=self.config.win_length,
452
+ power=2,
453
+ normalized=False,
454
+ center=self.config.center,
455
+ )
456
+ mel_spectrogram = (spectrogram.mT @ self.melscale_fbanks.to(torch.float32)).mT
457
+ # x has shape [batch, freq, time].
458
+ # F.amplitude_to_DB accepts inputs shaped as:
459
+ # - [freq, time]
460
+ # - [channel, freq, time]
461
+ # - [..., channel, freq, time]
462
+ # Here we insert a channel dimension of size 1 before calling it,
463
+ # then remove that extra dimension afterward.
464
+ log_mel_spectrogram = F.amplitude_to_DB(
465
+ mel_spectrogram.unsqueeze(1),
466
+ multiplier=10,
467
+ amin=1e-10,
468
+ db_multiplier=0,
469
+ top_db=120,
470
+ ).squeeze(1)
471
+ return log_mel_spectrogram.to(waveform.dtype)
472
+
473
+
474
+ class FixedAffine2d(nn.Module):
475
+ """
476
+ Per-channel fixed affine transform:
477
+ y = x * scale + bias
478
+ where scale/bias are broadcast on (B, C, H, W).
479
+ """
480
+
481
+ def __init__(self, scale: torch.Tensor, bias: torch.Tensor):
482
+ super().__init__()
483
+ self.register_buffer("scale", scale.reshape(1, -1, 1, 1))
484
+ self.register_buffer("bias", bias.reshape(1, -1, 1, 1))
485
+
486
+ @classmethod
487
+ def from_batchnorm2d(cls, bn: nn.BatchNorm2d) -> "FixedAffine2d":
488
+ if bn.running_mean is None or bn.running_var is None:
489
+ raise ValueError("BatchNorm2d must have running stats to be converted.")
490
+
491
+ if bn.affine:
492
+ gamma = bn.weight.detach()
493
+ beta = bn.bias.detach()
494
+ else:
495
+ gamma = torch.ones_like(bn.running_mean)
496
+ beta = torch.zeros_like(bn.running_mean)
497
+
498
+ running_mean = bn.running_mean.detach()
499
+ running_var = bn.running_var.detach()
500
+
501
+ scale = gamma / torch.sqrt(running_var + bn.eps)
502
+ bias = beta - running_mean * scale
503
+ return cls(scale=scale, bias=bias)
504
+
505
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
506
+ return x * self.scale + self.bias
507
+
508
+
509
+ class DashengAudioTransformer(PreTrainedModel):
510
+ config_class = DashengConfig
511
+ supports_gradient_checkpointing = True
512
+
513
+ def __init__(self, config: DashengConfig):
514
+ super().__init__(config)
515
+
516
+ self.target_length = config.target_length
517
+ self.embed_dim = config.embed_dim
518
+ self.hop_length = config.hop_length
519
+ self.gradient_checkpointing = False
520
+
521
+ self.front_end = DashengFrontend(config)
522
+
523
+ self.init_bn = nn.BatchNorm2d(config.n_mels, momentum=0.01)
524
+
525
+ self.patch_embed = AudioPatchEmbed(
526
+ input_size=(config.n_mels, config.target_length),
527
+ embed_dim=config.embed_dim,
528
+ in_chans=config.input_channels,
529
+ patch_size=config.patch_size,
530
+ flatten=False,
531
+ patch_stride=config.patch_stride,
532
+ )
533
+
534
+ self.time_pos_embed = nn.Parameter(
535
+ torch.randn(1, config.embed_dim, 1, self.patch_embed.grid_size[1]) * 0.02
536
+ )
537
+ self.freq_pos_embed = nn.Parameter(
538
+ torch.randn(1, config.embed_dim, self.patch_embed.grid_size[0], 1) * 0.02
539
+ )
540
+
541
+ self.pos_drop = nn.Dropout(p=config.drop_rate)
542
+ self.blocks = nn.ModuleList(
543
+ DashengBlock(
544
+ dim=config.embed_dim,
545
+ num_heads=config.num_heads,
546
+ mlp_ratio=config.mlp_ratio,
547
+ qkv_bias=config.qkv_bias,
548
+ init_values=config.init_values,
549
+ drop=config.drop_rate,
550
+ attn_drop=config.attn_drop_rate,
551
+ )
552
+ for _ in range(config.depth)
553
+ )
554
+ self.norm = nn.LayerNorm(config.embed_dim, eps=1e-6)
555
+
556
+ self.post_init()
557
+
558
+ def replace_init_bn_with_fixed_affine(self):
559
+ """
560
+ Call this after checkpoint is loaded and before inference/export.
561
+ """
562
+ if isinstance(self.init_bn, nn.BatchNorm2d):
563
+ self.init_bn.eval()
564
+ self.init_bn = FixedAffine2d.from_batchnorm2d(self.init_bn)
565
+
566
+ def forward_features(
567
+ self,
568
+ x: torch.Tensor,
569
+ mask: Optional[torch.Tensor] = None,
570
+ ) -> torch.Tensor:
571
+ t = x.shape[-1]
572
+ x = x + self.time_pos_embed[:, :, :, :t]
573
+ x = (
574
+ x + self.freq_pos_embed[:, :, :, :]
575
+ ) # Just to support __getitem__ in posembed
576
+ x = torch.permute(
577
+ torch.flatten(x, 2, 3), (0, 2, 1)
578
+ ) # rearrange(x, "b c f t -> b (f t) c")
579
+ x = self.pos_drop(x)
580
+ for block in self.blocks:
581
+ if self.gradient_checkpointing and self.training:
582
+ x = self._gradient_checkpointing_func(block, x, mask)
583
+ else:
584
+ x = block(x, mask)
585
+ x = self.norm(x)
586
+ return x
587
+
588
+ def _to_mask(self, lengths: torch.Tensor, max_length: int) -> torch.Tensor:
589
+ batch_size = len(lengths)
590
+ idx = torch.arange(max_length, device=lengths.device)
591
+ idx = idx.repeat(batch_size).view(batch_size, max_length)
592
+ mask = (idx < lengths.unsqueeze(-1)).bool()
593
+ return mask
594
+
595
+ def forward(
596
+ self,
597
+ x: torch.Tensor,
598
+ x_length: Optional[torch.Tensor] = None,
599
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
600
+ x = self.front_end(x)
601
+ target_length_in_patches = self.target_length // 4
602
+ x = x.unsqueeze(1)
603
+ x = torch.permute(x, (0, 2, 1, 3))
604
+ x = self.init_bn(x)
605
+ x = torch.permute(x, (0, 2, 1, 3))
606
+
607
+ x = self.patch_embed(x)
608
+ t = x.shape[-1]
609
+
610
+ input_splits = x.split(target_length_in_patches, dim=-1)
611
+
612
+ if x_length is not None:
613
+ assert len(x_length) == len(x), (
614
+ "batchsizes of input x and x_length need to be same"
615
+ )
616
+ assert x_length.ndim == 1, "Lengths are of size (B,)"
617
+ scaled_lengths = (x_length / (self.hop_length * 4)).long()
618
+ mask = self._to_mask(max_length=t, lengths=scaled_lengths)
619
+ split_masks = mask.split(target_length_in_patches, dim=-1)
620
+ else:
621
+ mask = None
622
+ split_masks = [None] * len(input_splits)
623
+
624
+ outputs = []
625
+
626
+ for split_x, split_mask in zip(input_splits, split_masks):
627
+ split_x = self.forward_features(split_x, mask=split_mask)
628
+ outputs.append(split_x)
629
+ x = torch.cat(outputs, dim=1)
630
+
631
+ return x, mask
632
+
633
+
634
+ class AudioProjectorSubsample(nn.Module):
635
+ def __init__(
636
+ self,
637
+ in_dim: int,
638
+ out_dim: int,
639
+ downsample_rate=5,
640
+ dtype: Optional[torch.dtype] = None,
641
+ ):
642
+ super().__init__()
643
+ self.k = downsample_rate
644
+ self.out_dim = out_dim
645
+ self.net = nn.Sequential(
646
+ nn.Linear(in_dim * self.k, out_dim, dtype=dtype),
647
+ nn.GELU(),
648
+ nn.Linear(out_dim, out_dim, dtype=dtype),
649
+ )
650
+
651
+ def forward(self, x, mask=None):
652
+ batch_size, seq_len, dim = x.shape
653
+ num_frames_to_discard = seq_len % self.k
654
+ if num_frames_to_discard > 0:
655
+ x = x[:, :-num_frames_to_discard, :]
656
+ if mask is not None:
657
+ mask = mask[:, :-num_frames_to_discard]
658
+ if mask is None:
659
+ mask = torch.ones(x.shape[:-1], dtype=torch.long, device=x.device)
660
+ x = x.reshape(
661
+ batch_size, -1, self.k * dim
662
+ ) # rearrange(x, "b (s k) d -> b s (k d)", k=self.k)
663
+ x = self.net(x)
664
+ mask = mask.reshape(
665
+ batch_size, -1, self.k
666
+ ) # rearrange(mask, "b (s k) -> b s k", k=self.k)
667
+ mask = mask.any(dim=-1).long()
668
+ return x, mask
669
+
670
+
671
+
672
+ @dataclass
673
+ class Qwen25OmniTextModelOutput(ModelOutput):
674
+ loss: Optional[torch.FloatTensor] = None
675
+ logits: Optional[torch.FloatTensor] = None
676
+ past_key_values: Optional[Cache] = None
677
+ hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
678
+ attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
679
+
680
+
681
+ class Qwen25OmniThinkerTextOnlyDecoder(PreTrainedModel, GenerationMixin):
682
+ config_class = Qwen2_5OmniTextConfig
683
+ _supports_flash_attn_2 = True
684
+ _supports_sdpa = True
685
+ _supports_cache_class = True
686
+ _supports_static_cache = True
687
+
688
+ def __init__(self, config: Qwen2_5OmniTextConfig):
689
+ super().__init__(config)
690
+ self.model = Qwen2_5OmniThinkerTextModel._from_config(config)
691
+ self.lm_head = nn.Linear(
692
+ config.hidden_size,
693
+ config.vocab_size,
694
+ bias=False,
695
+ )
696
+ self.post_init()
697
+
698
+ @can_return_tuple
699
+ def forward(
700
+ self,
701
+ input_ids: Optional[torch.LongTensor] = None,
702
+ attention_mask: Optional[torch.Tensor] = None,
703
+ position_ids: Optional[torch.LongTensor] = None,
704
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
705
+ inputs_embeds: Optional[torch.FloatTensor] = None,
706
+ use_cache: Optional[bool] = None,
707
+ output_attentions: Optional[bool] = None,
708
+ output_hidden_states: Optional[bool] = None,
709
+ cache_position: Optional[torch.LongTensor] = None,
710
+ labels: Optional[torch.Tensor] = None,
711
+ **kwargs,
712
+ ) -> Union[Tuple, Qwen25OmniTextModelOutput]:
713
+ if attention_mask is not None and position_ids is None:
714
+ position_ids = (
715
+ attention_mask.long()
716
+ .cumsum(dim=-1)
717
+ .masked_fill_(attention_mask == 0, 1)
718
+ - 1
719
+ )
720
+
721
+ outputs: BaseModelOutputWithPast = self.model(
722
+ input_ids=input_ids,
723
+ attention_mask=attention_mask,
724
+ position_ids=position_ids,
725
+ past_key_values=past_key_values,
726
+ inputs_embeds=inputs_embeds,
727
+ use_cache=use_cache,
728
+ output_attentions=output_attentions,
729
+ output_hidden_states=output_hidden_states,
730
+ cache_position=cache_position,
731
+ return_dict=True,
732
+ )
733
+ hidden_states = outputs.last_hidden_state
734
+ logits = self.lm_head(hidden_states)
735
+
736
+ loss = (
737
+ self.loss_function(
738
+ logits=logits,
739
+ labels=labels,
740
+ vocab_size=self.config.vocab_size,
741
+ **kwargs,
742
+ )
743
+ if labels is not None
744
+ else None
745
+ )
746
+
747
+ return Qwen25OmniTextModelOutput(
748
+ loss=loss,
749
+ logits=logits,
750
+ past_key_values=outputs.past_key_values,
751
+ hidden_states=outputs.hidden_states,
752
+ attentions=outputs.attentions,
753
+ )
754
+
755
+
756
+ # Hardcoded architecture / training choices (not exposed in config.json).
757
+ USE_LOGIT_SCALE = True
758
+ HIDDEN_SIZE = 3584
759
+
760
+
761
+ class ALM2VecModel(PreTrainedModel):
762
+ config_class = ALM2VecConfig
763
+
764
+ def __init__(self, config: ALM2VecConfig):
765
+ super().__init__(config)
766
+ text_config = Qwen2_5OmniTextConfig(**DASHENG_ARCH_CONFIG["text_config"])
767
+ decoder = Qwen25OmniThinkerTextOnlyDecoder(text_config)
768
+ self.model = decoder.model
769
+ self.dasheng = DashengAudioTransformer(
770
+ DashengConfig(**DASHENG_ARCH_CONFIG["audio_encoder_config"])
771
+ )
772
+ self.dasheng_down = AudioProjectorSubsample(**DASHENG_ARCH_CONFIG["audio_projector_config"])
773
+ self.dasheng_proj = nn.Identity()
774
+
775
+ # Placeholder shapes match exported checkpoints (overwritten by load_state_dict).
776
+ self.register_buffer("audio_start_token", torch.zeros(1, dtype=torch.long))
777
+ self.register_buffer("audio_end_token", torch.zeros(1, dtype=torch.long))
778
+ self.register_buffer("eos_token", torch.zeros(7, dtype=torch.long))
779
+
780
+ self.hidden_size = self.model.config.hidden_size
781
+ self.use_checkpointing = False
782
+ self.checkpoint_reentrant = False
783
+
784
+ if USE_LOGIT_SCALE:
785
+ init_value = math.log(1 / 0.07)
786
+ self.register_buffer(
787
+ "logit_scale", torch.tensor([init_value], dtype=torch.float32)
788
+ )
789
+ else:
790
+ self.logit_scale = None
791
+
792
+ self.siglip_head = nn.Linear(self.hidden_size, self.hidden_size)
793
+ self.dasheng.replace_init_bn_with_fixed_affine()
794
+ self._tokenizer = None
795
+ self.post_init()
796
+
797
+ def set_tokenizer(self, tokenizer) -> None:
798
+ """Attach a tokenizer for high-level encode APIs."""
799
+ self._tokenizer = tokenizer
800
+
801
+ def _resolve_tokenizer(self):
802
+ if self._tokenizer is not None:
803
+ return self._tokenizer
804
+ tokenizer = AutoTokenizer.from_pretrained(self.name_or_path, trust_remote_code=True)
805
+ self._tokenizer = tokenizer
806
+ return tokenizer
807
+
808
+ @staticmethod
809
+ def _to_list(x: Any) -> list[Any]:
810
+ if x is None:
811
+ return []
812
+ if isinstance(x, (list, tuple)):
813
+ return list(x)
814
+ return [x]
815
+
816
+ @staticmethod
817
+ def _build_prompt(instruction: str, text: Optional[str]) -> str:
818
+ system_part = f"<|im_start|>system\n{instruction}<|im_end|>\n"
819
+ if text is None:
820
+ user_part = "<|im_start|>user\n"
821
+ else:
822
+ user_part = f"<|im_start|>user\n{text}"
823
+ return system_part + user_part
824
+
825
+ def _prepare_text_batch(
826
+ self,
827
+ tokenizer,
828
+ texts: list[Optional[str]],
829
+ *,
830
+ task: str,
831
+ instruction: Optional[str],
832
+ device: torch.device,
833
+ ) -> tuple[torch.Tensor, torch.Tensor]:
834
+ if task not in ("query", "document"):
835
+ raise ValueError(f"Unsupported task={task}. Use 'query' or 'document'.")
836
+ default_instruction = QUERY_INSTRUCTION if task == "query" else DOC_INSTRUCTION
837
+ instruction = instruction or default_instruction
838
+ prompts = [self._build_prompt(instruction, text) for text in texts]
839
+ encoded = tokenizer(prompts, padding=True, add_special_tokens=False, return_tensors="pt")
840
+ text_ids = encoded["input_ids"].to(device)
841
+ text_lens = encoded["attention_mask"].to(device).sum(dim=1)
842
+ return text_ids, text_lens
843
+
844
+ def _load_audio_path(self, path: Union[str, Path], target_sr: int) -> torch.Tensor:
845
+ raw_path = str(path)
846
+ parsed = urlparse(raw_path)
847
+ is_remote_url = parsed.scheme in ("http", "https")
848
+ local_path = None if is_remote_url else Path(path)
849
+ suffix_source = Path(parsed.path) if is_remote_url else local_path
850
+ suffix = suffix_source.suffix.lower()
851
+
852
+ if is_remote_url:
853
+ with urlopen(raw_path) as resp:
854
+ audio_bytes = resp.read()
855
+ source_for_torchaudio = NamedTemporaryFile(
856
+ suffix=suffix or ".audio",
857
+ delete=False,
858
+ )
859
+ source_for_torchaudio.write(audio_bytes)
860
+ source_for_torchaudio.flush()
861
+ source_for_torchaudio.close()
862
+ else:
863
+ source_for_torchaudio = None
864
+
865
+ path_for_wave = BytesIO(audio_bytes) if is_remote_url else str(local_path)
866
+
867
+ try:
868
+ if suffix in (".wav", ".wave"):
869
+ with wave.open(path_for_wave, "rb") as wf:
870
+ sr = wf.getframerate()
871
+ n_channels = wf.getnchannels()
872
+ sample_width = wf.getsampwidth()
873
+ raw = wf.readframes(wf.getnframes())
874
+ if sample_width == 1:
875
+ audio = torch.frombuffer(bytearray(raw), dtype=torch.uint8).float()
876
+ audio = (audio - 128.0) / 128.0
877
+ elif sample_width == 2:
878
+ audio = torch.frombuffer(bytearray(raw), dtype=torch.int16).float() / 32768.0
879
+ elif sample_width == 4:
880
+ audio = torch.frombuffer(bytearray(raw), dtype=torch.int32).float() / 2147483648.0
881
+ else:
882
+ raise ValueError(f"Unsupported WAV sample width: {sample_width}")
883
+ if n_channels > 1:
884
+ audio = audio.reshape(-1, n_channels).mean(dim=1)
885
+ else:
886
+ if torchaudio is None:
887
+ raise ImportError("torchaudio is required for non-WAV audio paths.")
888
+ load_target = (
889
+ source_for_torchaudio.name if is_remote_url else str(local_path)
890
+ )
891
+ waveform, sr = torchaudio.load(load_target)
892
+ if waveform.shape[0] > 1:
893
+ waveform = waveform.mean(dim=0, keepdim=True)
894
+ audio = waveform.squeeze(0)
895
+ finally:
896
+ if source_for_torchaudio is not None:
897
+ try:
898
+ Path(source_for_torchaudio.name).unlink(missing_ok=True)
899
+ except OSError:
900
+ pass
901
+ if sr != target_sr:
902
+ if torchaudio is None:
903
+ raise ImportError("torchaudio is required for resampling.")
904
+ audio = torchaudio.functional.resample(audio.unsqueeze(0), sr, target_sr).squeeze(0)
905
+ return audio.float()
906
+
907
+ def _prepare_audio_batch(
908
+ self,
909
+ audio_items: list[Optional[Union[str, Path, torch.Tensor]]],
910
+ *,
911
+ target_sr: int,
912
+ device: torch.device,
913
+ ) -> tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
914
+ tensor_list: list[torch.Tensor] = []
915
+ lens_list: list[int] = []
916
+ has_audio = False
917
+ for item in audio_items:
918
+ if item is None:
919
+ tensor_list.append(torch.zeros(1, dtype=torch.float32))
920
+ lens_list.append(0)
921
+ continue
922
+ has_audio = True
923
+ if isinstance(item, (str, Path)):
924
+ wav = self._load_audio_path(item, target_sr=target_sr)
925
+ elif isinstance(item, torch.Tensor):
926
+ wav = item.detach().float().cpu()
927
+ if wav.dim() == 2:
928
+ wav = wav.mean(dim=0)
929
+ elif wav.dim() != 1:
930
+ raise ValueError("Audio tensor must be 1D waveform or 2D [channels, length].")
931
+ else:
932
+ raise TypeError(f"Unsupported audio item type: {type(item)}")
933
+ if wav.numel() == 0:
934
+ wav = torch.zeros(1, dtype=torch.float32)
935
+ length = 0
936
+ else:
937
+ length = int(wav.numel())
938
+ tensor_list.append(wav)
939
+ lens_list.append(length)
940
+ if not has_audio:
941
+ return None, None
942
+ max_len = max(t.numel() for t in tensor_list)
943
+ padded = torch.zeros(len(tensor_list), max_len, dtype=torch.float32)
944
+ for i, wav in enumerate(tensor_list):
945
+ L = wav.numel()
946
+ if L > 0:
947
+ padded[i, :L] = wav
948
+ return padded.to(device), torch.tensor(lens_list, dtype=torch.long, device=device)
949
+
950
+ def encode(
951
+ self,
952
+ *,
953
+ text: Optional[Union[str, list[str]]] = None,
954
+ audio: Optional[Union[str, Path, torch.Tensor, list[Optional[Union[str, Path, torch.Tensor]]]]] = None,
955
+ task: str = "document",
956
+ instruction: Optional[str] = None,
957
+ normalize: bool = True,
958
+ device: Optional[Union[str, torch.device]] = None,
959
+ ) -> torch.Tensor:
960
+ """High-level embedding API. Accepts raw text/audio and returns embeddings."""
961
+ self.eval()
962
+ tokenizer = self._resolve_tokenizer()
963
+ device = torch.device(device) if device is not None else next(self.parameters()).device
964
+
965
+ text_items = self._to_list(text)
966
+ audio_items = self._to_list(audio)
967
+ batch_size = max(len(text_items), len(audio_items))
968
+ if batch_size == 0:
969
+ raise ValueError("At least one of text/audio must be provided.")
970
+
971
+ if len(text_items) == 0:
972
+ text_items = [None] * batch_size
973
+ elif len(text_items) == 1 and batch_size > 1:
974
+ text_items = text_items * batch_size
975
+ elif len(text_items) != batch_size:
976
+ raise ValueError("text and audio batch sizes must match (or be broadcastable length 1).")
977
+
978
+ if len(audio_items) == 0:
979
+ audio_items = [None] * batch_size
980
+ elif len(audio_items) == 1 and batch_size > 1:
981
+ audio_items = audio_items * batch_size
982
+ elif len(audio_items) != batch_size:
983
+ raise ValueError("text and audio batch sizes must match (or be broadcastable length 1).")
984
+
985
+ text_ids, text_lens = self._prepare_text_batch(
986
+ tokenizer,
987
+ texts=text_items,
988
+ task=task,
989
+ instruction=instruction,
990
+ device=device,
991
+ )
992
+ audio_tensor, audio_lens = self._prepare_audio_batch(
993
+ audio_items,
994
+ target_sr=TARGET_SR,
995
+ device=device,
996
+ )
997
+
998
+ with torch.inference_mode():
999
+ emb = self(
1000
+ text_ids=text_ids,
1001
+ text_lens=text_lens,
1002
+ audio=audio_tensor,
1003
+ audio_lens=audio_lens,
1004
+ )
1005
+ if normalize:
1006
+ emb = torch.nn.functional.normalize(emb.float(), dim=-1)
1007
+ return emb
1008
+
1009
+ @staticmethod
1010
+ def _check_list_arg(name: str, value: Any, item_types: Tuple[type, ...]) -> None:
1011
+ if value is None:
1012
+ return
1013
+ if not isinstance(value, list):
1014
+ raise TypeError(
1015
+ f"`{name}` must be a list, got {type(value).__name__}."
1016
+ )
1017
+ if len(value) == 0:
1018
+ raise ValueError(f"`{name}` must be a non-empty list when provided.")
1019
+ for i, item in enumerate(value):
1020
+ if not isinstance(item, item_types):
1021
+ allowed = ", ".join(t.__name__ for t in item_types)
1022
+ raise TypeError(
1023
+ f"`{name}[{i}]` must be one of ({allowed}), "
1024
+ f"got {type(item).__name__}."
1025
+ )
1026
+
1027
+ def encode_query(
1028
+ self,
1029
+ text: Optional[List[str]] = None,
1030
+ audio: Optional[List[Union[str, Path, torch.Tensor]]] = None,
1031
+ **kwargs,
1032
+ ) -> torch.Tensor:
1033
+ """Encode queries. Accepts text-only, audio-only, or both."""
1034
+ self._check_list_arg("text", text, (str,))
1035
+ self._check_list_arg("audio", audio, (str, Path, torch.Tensor))
1036
+ if text is None and audio is None:
1037
+ raise ValueError(
1038
+ "encode_query requires at least one of `text` or `audio`."
1039
+ )
1040
+ if text is not None and audio is not None and len(text) != len(audio):
1041
+ raise ValueError(
1042
+ f"encode_query: `text` (len={len(text)}) and `audio` (len={len(audio)}) "
1043
+ "must have the same length when both are provided."
1044
+ )
1045
+ return self.encode(text=text, audio=audio, task="query", **kwargs)
1046
+
1047
+ def encode_document(
1048
+ self,
1049
+ text: Optional[List[str]] = None,
1050
+ audio: Optional[List[Union[str, Path, torch.Tensor]]] = None,
1051
+ **kwargs,
1052
+ ) -> torch.Tensor:
1053
+ """Encode documents. Accepts exactly one of `text` or `audio`."""
1054
+ self._check_list_arg("text", text, (str,))
1055
+ self._check_list_arg("audio", audio, (str, Path, torch.Tensor))
1056
+ if (text is None) == (audio is None):
1057
+ raise ValueError(
1058
+ "encode_document requires exactly one of `text` or `audio` "
1059
+ "(not both, not neither)."
1060
+ )
1061
+ return self.encode(text=text, audio=audio, task="document", **kwargs)
1062
+
1063
+ def forward(self, text_ids, text_lens, audio=None, audio_lens=None):
1064
+ device = text_ids.device
1065
+ B = text_ids.size(0)
1066
+ embed_dtype = self.model.embed_tokens.weight.dtype
1067
+
1068
+ if audio is not None:
1069
+ audio_emb, audio_mask = self.dasheng(
1070
+ audio,
1071
+ audio_lens,
1072
+ )
1073
+ audio_emb, audio_mask = self.dasheng_down(audio_emb, audio_mask)
1074
+ audio_lens = audio_mask.sum(dim=1)
1075
+ audio_emb = self.dasheng_proj(audio_emb)
1076
+ else:
1077
+ audio_emb = None
1078
+
1079
+ text_emb = self.model.embed_tokens(text_ids)
1080
+ audio_start_emb = self.model.embed_tokens(self.audio_start_token.clone())
1081
+ audio_end_emb = self.model.embed_tokens(self.audio_end_token.clone())
1082
+ eos_emb = self.model.embed_tokens(self.eos_token.clone())
1083
+
1084
+ input_embeds = []
1085
+ attention_masks = []
1086
+ last_indices = []
1087
+
1088
+ for i in range(B):
1089
+ seq = [text_emb[i, : text_lens[i]]]
1090
+ if audio_emb is not None and audio_lens[i] > 0:
1091
+ seq.append(audio_start_emb)
1092
+ seq.append(audio_emb[i, : audio_lens[i]])
1093
+ seq.append(audio_end_emb)
1094
+ seq.append(eos_emb)
1095
+ seq = torch.cat(seq, dim=0)
1096
+ input_embeds.append(seq)
1097
+ attention_masks.append(torch.ones(seq.size(0), device=device))
1098
+ last_indices.append(seq.size(0) - 1)
1099
+
1100
+ max_len = max(x.size(0) for x in input_embeds)
1101
+ padded_embeds = torch.zeros(
1102
+ B, max_len, self.hidden_size, device=device, dtype=embed_dtype
1103
+ )
1104
+ padded_mask = torch.zeros(B, max_len, device=device)
1105
+ for i in range(B):
1106
+ L = input_embeds[i].size(0)
1107
+ padded_embeds[i, :L] = input_embeds[i]
1108
+ padded_mask[i, :L] = attention_masks[i]
1109
+
1110
+ outputs = self.model(
1111
+ inputs_embeds=padded_embeds,
1112
+ attention_mask=padded_mask,
1113
+ ).last_hidden_state
1114
+
1115
+ final_hidden = torch.stack([outputs[i, last_indices[i]] for i in range(B)])
1116
+ out = self.siglip_head(final_hidden).squeeze(1)
1117
+ return out
1118
+