File size: 15,849 Bytes
45cf443 c6bc767 45cf443 92c7321 45cf443 c6bc767 45cf443 7036629 45cf443 fe5285c 7036629 fe5285c 7036629 45cf443 fe5285c 45cf443 fe5285c 45cf443 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 | import io
import json
import random
import torch
from torch.utils.data import Dataset
from PIL import Image
from dataset.common import pre_processing_chat, post_processing_chat
class VAMDataset(Dataset):
def __init__(self, data_path, tokenizer, audio_processor=None, vision_processor=None,
max_length=1200, audio_special_token='<|audio_pad|>', image_special_token='<|image_pad|>',
audio_stop_token=2050, # <|audio_stop|>
audio_pad_token=2049, # <|audio_pad|>
audio_spk_token=2051, # <|audio_spk|>
audio_vocab_size=2112, # 2048 mimi codes + 64 special tokens
scheduled_sampling=0.05,
image_token_len=64, max_samples=None):
super().__init__()
import pyarrow as pa
import pyarrow.parquet as pq
tables = []
total = 0
for p in data_path.split(','):
pf = pq.ParquetFile(p.strip())
for batch in pf.iter_batches(batch_size=4096):
tables.append(batch)
total += batch.num_rows
if max_samples is not None and total >= max_samples:
break
if max_samples is not None and total >= max_samples:
break
if tables:
self.table = pa.Table.from_batches(tables)
self.table = self.table.cast(pa.schema([f.with_type(pa.large_string()) if pa.types.is_string(f.type) else f for f in self.table.schema]))
if max_samples is not None:
self.table = self.table.slice(0, min(max_samples, len(self.table)))
else:
self.table = pa.Table.from_pydict({})
self.tokenizer = tokenizer
self.audio_processor = audio_processor
self.vision_processor = vision_processor
self.max_length = max_length
self.audio_token = audio_special_token
self.image_token_len = image_token_len
self.image_token = image_special_token * image_token_len
self.audio_stop_token = audio_stop_token
self.audio_pad_token = audio_pad_token
self.audio_spk_token = audio_spk_token
self.audio_vocab_size = audio_vocab_size
self.scheduled_sampling_prob = scheduled_sampling
self.text_vocab_size = len(tokenizer)
self.image_token_id = tokenizer.encode(image_special_token, add_special_tokens=False)[0]
self.audio_token_id = tokenizer.encode(audio_special_token, add_special_tokens=False)[0]
self.think_end_ids = tokenizer.encode('</think>\n\n', add_special_tokens=False)
self.bos_id = tokenizer(f'{tokenizer.bos_token}assistant\n', add_special_tokens=False).input_ids
self.eos_id = tokenizer(f'{tokenizer.eos_token}\n', add_special_tokens=False).input_ids
def __len__(self):
return len(self.table)
@staticmethod
def process_audio(audio_path, audio_processor):
import soundfile as sf
import numpy as np
wav, sr = sf.read(audio_path)
if wav.ndim > 1: wav = wav.mean(axis=1)
if sr != 16000:
import librosa
wav = librosa.resample(wav.astype(float), orig_sr=sr, target_sr=16000)
inputs = audio_processor(wav.astype(np.float32), sampling_rate=16000, return_tensors="pt", return_attention_mask=True)
valid_len = inputs.attention_mask.sum().item()
return inputs.input_features.squeeze(0), valid_len
def augment_wav(self, wav, sr=16000):
import numpy as np
from scipy.signal import resample
if random.random() < 0.5:
speed = random.uniform(0.7, 1.6)
wav = resample(wav, int(len(wav) / speed)).astype(np.float32)
if random.random() < 0.3:
noise = np.random.randn(len(wav)).astype(np.float32) * random.uniform(0.001, 0.01)
wav = wav + noise
if random.random() < 0.3:
wav = wav * random.uniform(0.8, 1.2)
if random.random() < 0.2 and len(wav) > sr:
start = random.randint(0, len(wav) - sr // 4)
wav[start:start + sr // 4] = 0
if random.random() < 0.2:
k = random.choice([3, 5, 7])
wav = np.convolve(wav, np.ones(k) / k, mode='same').astype(np.float32)
if random.random() < 0.3:
ir_len = int(sr * random.uniform(0.05, 0.2))
ir = np.random.randn(ir_len).astype(np.float32) * np.exp(-np.linspace(0, 10, ir_len))
ir[0] = 1.0
ir /= np.sqrt(np.sum(ir ** 2) + 1e-6)
wav = np.convolve(wav, ir, mode='same').astype(np.float32)
if random.random() < 0.2:
pink = np.cumsum(np.random.randn(len(wav))).astype(np.float32)
pink /= np.max(np.abs(pink)) + 1e-6
wav = wav + pink * random.uniform(0.003, 0.015)
return np.clip(wav, -1.0, 1.0).astype(np.float32)
def augment_mel(self, fbank):
import numpy as np
T, D = fbank.shape
if random.random() < 0.5:
f = random.randint(1, 64)
f0 = random.randint(0, D - f)
fbank[:, f0:f0 + f] = 0
if random.random() < 0.5 and T > 1:
t = random.randint(1, min(10, T))
t0 = random.randint(0, T - t)
fbank[t0:t0 + t, :] = 0
return fbank
def load_audio_inputs(self, audio_bytes):
import soundfile as sf
import numpy as np
import io
import torch
if not audio_bytes: return None, 0
wav, sr = sf.read(io.BytesIO(audio_bytes))
if wav.ndim > 1: wav = wav.mean(axis=1)
wav = wav.astype(np.float32)
if sr != 16000:
import torchaudio.functional as AF
wav_t = torch.from_numpy(wav).unsqueeze(0)
wav_t = AF.resample(wav_t, sr, 16000)
wav = wav_t.squeeze(0).numpy()
wav = self.augment_wav(wav)
inputs = self.audio_processor(wav, sampling_rate=16000, return_tensors="pt", return_attention_mask=True)
valid_len = inputs.attention_mask.sum().item()
return self.augment_mel(inputs.input_features.squeeze(0)), valid_len
def load_image_inputs(self, image_bytes):
import io
from PIL import Image
if not image_bytes or self.vision_processor is None: return None
image = Image.open(io.BytesIO(image_bytes)).convert('RGB')
inputs = self.vision_processor(images=image, return_tensors="pt")
if hasattr(inputs, 'keys'): return {k: v for k, v in inputs.items()}
return inputs.pixel_values
def create_chat_prompt(self, conversations, audio_features_length=0):
conversations = pre_processing_chat(conversations)
messages = []
is_last_user = lambda i: i == max(j for j, t in enumerate(conversations) if t['role'] == 'user')
for idx, turn in enumerate(conversations):
role, content = turn['role'], turn['content']
if role == 'user' and is_last_user(idx) and audio_features_length > 0:
ap = self.audio_token * audio_features_length
r = random.random()
if r < 0.4: content = ap
elif r < 0.6: content = content
elif r < 0.8: content = ap + '\n\n' + content
else: content = content + '\n\n' + ap
if '<image>' in content:
r = random.random()
if r < 0.2: content = '<image>\n' + content.replace('<image>', '').strip()
elif r < 0.4: content = '<image>\n\n' + content.replace('<image>', '').strip()
elif r < 0.6: content = content.replace('<image>', '').strip() + '\n' + '<image>'
else: content = content.replace('<image>', '').strip() + '\n\n' + '<image>'
messages.append({"role": role, "content": content})
prompt = self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=False)
return post_processing_chat(prompt)
def generate_text_labels(self, input_ids):
labels = [-100] * len(input_ids)
ranges = []
i = 0
while i < len(input_ids):
if input_ids[i:i + len(self.bos_id)] == self.bos_id:
start = i + len(self.bos_id)
end = start
while end < len(input_ids):
if input_ids[end:end + len(self.eos_id)] == self.eos_id:
break
end += 1
ranges.append((start, end))
for j in range(start, min(end + len(self.eos_id), self.max_length)):
labels[j] = input_ids[j]
i = end + len(self.eos_id) if end < len(input_ids) else len(input_ids)
else:
i += 1
return labels, ranges
def apply_scheduled_sampling(self, input_ids, audio_labels, text_labels):
if self.scheduled_sampling_prob <= 0:
return input_ids
audio_mask = (audio_labels != -100).any(dim=0) & (torch.rand(input_ids.size(1)) < self.scheduled_sampling_prob)
for i in range(8):
input_ids[i] = torch.where(audio_mask, torch.randint(0, self.audio_vocab_size, input_ids[i].shape), input_ids[i])
text_mask = (text_labels != -100) & (input_ids[8] != self.image_token_id) & (torch.rand(input_ids.size(1)) < self.scheduled_sampling_prob)
input_ids[8] = torch.where(text_mask, torch.randint(0, self.text_vocab_size, input_ids[8].shape), input_ids[8])
return input_ids
def __getitem__(self, index: int):
import numpy as np
conversations = json.loads(self.table['conversations'][index].as_py())
question_audios = self.table['question_audios'][index].as_py() if 'question_audios' in self.table.column_names else []
answer_audios = self.table['answer_audios'][index].as_py() if 'answer_audios' in self.table.column_names else []
image_bytes = self.table['image_bytes'][index].as_py() if 'image_bytes' in self.table.column_names else []
if image_bytes and not isinstance(image_bytes, list): image_bytes = [image_bytes]
ref_audios = self.table['ref_audios'][index].as_py() if 'ref_audios' in self.table.column_names else []
spk_emb_raw = self.table['spk_emb'][index].as_py() if 'spk_emb' in self.table.column_names else []
asst_indices = [i for i, t in enumerate(conversations) if t['role'] == 'assistant']
if len(asst_indices) > 1:
rand_idx = random.randint(0, len(asst_indices) - 1)
for i in range(rand_idx, -1, -1):
conversations = conversations[:asst_indices[i] + 1]
test_prompt = self.create_chat_prompt(conversations, 0)
if len(self.tokenizer(test_prompt).input_ids) + 100 < self.max_length:
break
pixel_values = None
user_count = sum(1 for t in conversations if t['role'] == 'user')
if image_bytes and len(image_bytes) > 0 and self.vision_processor:
pixel_values = self.load_image_inputs(image_bytes[0])
audio_inputs, audio_len, audio_features_length = None, 0, 0
user_count = sum(1 for t in conversations if t['role'] == 'user')
if question_audios and user_count > 0 and user_count <= len(question_audios) and self.audio_processor:
audio_bytes = question_audios[user_count - 1]
if audio_bytes:
mel, valid_len = self.load_audio_inputs(audio_bytes)
if mel is not None:
audio_inputs = mel.unsqueeze(0)
audio_len = valid_len
audio_features_length = valid_len or 1
if audio_inputs is None and self.audio_processor:
audio_inputs = torch.zeros(1, 1, 560)
audio_len = 0
if pixel_values is None and self.vision_processor:
pixel_values = {'pixel_values': torch.zeros(1, 3, 256, 256)}
last_audio_codes = None
asst_count = sum(1 for t in conversations if t['role'] == 'assistant')
if answer_audios and asst_count > 0 and asst_count <= len(answer_audios):
tokens = answer_audios[asst_count - 1]
if tokens:
audio_codes_8layers = [[] for _ in range(8)]
for i in range(0, len(tokens) - 7, 8):
for j in range(8): audio_codes_8layers[j].append(tokens[i + j])
for layer in audio_codes_8layers: layer.append(self.audio_stop_token)
last_audio_codes = audio_codes_8layers
prompt = self.create_chat_prompt(conversations, audio_features_length)
if pixel_values is not None: prompt = prompt.replace('<image>', self.image_token)
input_ids = self.tokenizer(prompt).input_ids[:self.max_length]
input_ids += [self.tokenizer.pad_token_id] * (self.max_length - len(input_ids))
text_labels, assistant_ranges = self.generate_text_labels(input_ids)
for start, end in assistant_ranges[:-1]:
mask_end = min(end + len(self.eos_id), self.max_length)
text_labels[start:mask_end] = [-100] * (mask_end - start)
Y_audio_layers = [[self.audio_pad_token] * self.max_length for _ in range(8)]
audio_labels = [[-100] * self.max_length for _ in range(8)]
if assistant_ranges and last_audio_codes:
assistant_start, assistant_end = assistant_ranges[-1]
for pos in range(assistant_start, min(assistant_end, assistant_start + 50)):
if input_ids[pos:pos + len(self.think_end_ids)] == self.think_end_ids:
assistant_start = pos + len(self.think_end_ids)
break
has_spk = bool(spk_emb_raw)
has_ref = bool(ref_audios) and random.random() > 0.5
spk_reserve = 1 if has_spk else 0
if has_ref:
ref_codes = [[] for _ in range(8)]
for i in range(0, len(ref_audios) - 7, 8):
for j in range(8): ref_codes[j].append(ref_audios[i + j])
ref_len = len(ref_codes[0])
ref_start = max(spk_reserve, assistant_start - ref_len)
for layer_idx in range(8):
codes = ref_codes[layer_idx][-(assistant_start - ref_start):] if ref_len > (assistant_start - ref_start) else ref_codes[layer_idx]
for i, code in enumerate(codes):
Y_audio_layers[layer_idx][ref_start + i] = code
else:
ref_start = assistant_start
if has_spk and ref_start > 0:
spk_pos = ref_start - 1
for layer_idx in range(8):
Y_audio_layers[layer_idx][spk_pos] = self.audio_spk_token
for layer_idx in range(8):
codes = last_audio_codes[layer_idx]
start_pos = assistant_start + layer_idx + 1
for i, code in enumerate(codes):
if start_pos + i < self.max_length:
Y_audio_layers[layer_idx][start_pos + i] = code
audio_labels[layer_idx][start_pos + i] = code
X_audio = torch.tensor([layer[:-1] for layer in Y_audio_layers], dtype=torch.long) # (8, T-1)
X_text = torch.tensor(input_ids[:-1], dtype=torch.long) # (T-1,)
input_ids = torch.cat((X_audio, X_text.unsqueeze(0)), dim=0) # (9, T-1)
text_labels = torch.tensor(text_labels[1:], dtype=torch.long) # (T-1,)
audio_labels = torch.tensor([layer[1:] for layer in audio_labels], dtype=torch.long) # (8, T-1)
input_ids = self.apply_scheduled_sampling(input_ids, audio_labels, text_labels)
spk_emb = torch.tensor(spk_emb_raw, dtype=torch.float32) if spk_emb_raw else torch.zeros(192)
return input_ids, text_labels, audio_labels, audio_inputs, audio_len, pixel_values, spk_emb
|