Files
soprano-factory/training/collator.py
T
2026-02-09 01:01:44 -05:00

41 lines
1.7 KiB
Python

# training/collator.py
import torch
import random
class SopranoCollator:
def __init__(self, tokenizer, seq_len=1024):
self.tokenizer = tokenizer
self.seq_len = seq_len
self.pad_token_id = tokenizer.pad_token_id or tokenizer.eos_token_id
def pack_for_llm(self, batch):
texts = [item["text"] for item in batch]
encodings = self.tokenizer(texts, add_special_tokens=False, padding=False, truncation=False)
input_ids_list = encodings["input_ids"]
packed_batch = []
buffer = []
buffer_len = 0
random.shuffle(input_ids_list)
for ids in input_ids_list:
ids = torch.tensor(ids, dtype=torch.long)
if buffer_len + len(ids) > self.seq_len:
full_seq = torch.cat(buffer)
if len(full_seq) < self.seq_len + 1:
padding = torch.full((self.seq_len + 1 - len(full_seq),), self.pad_token_id, dtype=torch.long)
full_seq = torch.cat([full_seq, padding])
packed_batch.append(full_seq[:self.seq_len + 1])
buffer, buffer_len = [], 0
buffer.append(ids)
buffer_len += len(ids)
if not packed_batch: return None, None
batch_tensor = torch.stack(packed_batch)
return batch_tensor[:, :-1], batch_tensor[:, 1:]
def collate_for_decoder(self, batch):
texts = [item["text"] for item in batch]
wav_paths = [item["wav_path"] for item in batch]
encodings = self.tokenizer(texts, padding=True, truncation=True, max_length=self.seq_len, return_tensors="pt", add_special_tokens=False)
return encodings["input_ids"], wav_paths