import gzip import pickle import torch import yaml from torch.utils.data import Dataset from transformers import AutoTokenizer, Trainer, TrainingArguments from src.configuration import SLTConfig from src.model import SLTModel def load_data(sgn_path): f = gzip.open(sgn_path, "rb") folders = pickle.load(f) return folders class SignLanguageDataset(Dataset): def __init__(self, pickle_path): data = load_data(pickle_path) self.signs = [torch.tensor(item["sign"], dtype=torch.float32) for item in data] self.texts = [item["text"] for item in data] def __len__(self): return len(self.signs) def __getitem__(self, idx): return { "sign": self.signs[idx], # [seq_len, feature_dim] "text": self.texts[idx], # str } class SLTDataCollator: def __init__(self, tokenizer, max_len=128): self.tokenizer = tokenizer self.max_len = max_len def __call__(self, batch): # 1. Process Signs (Padding) # List of tensors [seq_len, dim] -> padded batch [batch, max_seq, dim] sign_features = [item["sign"] for item in batch] # Pad sequence padded_signs = torch.nn.utils.rnn.pad_sequence(sign_features, batch_first=True) # Create Attention Mask for signs (1 for real, 0 for pad) batch_size = len(sign_features) max_sign_len = padded_signs.size(1) attention_mask = torch.zeros(batch_size, max_sign_len, dtype=torch.long) for i, sign in enumerate(sign_features): length = sign.size(0) attention_mask[i, :length] = 1 # 2. Process Text (Tokenization) texts = [item["text"] for item in batch] labels = self.tokenizer( texts, padding=True, truncation=True, return_tensors="pt", max_length=self.max_len, ).input_ids # Replace padding token id with -100 for loss calculation ignore labels[labels == self.tokenizer.pad_token_id] = -100 return { "input_values": padded_signs, "attention_mask": attention_mask, "labels": labels, } def main(): # load config file with open("configs/config.yaml") as f: try: config_file = yaml.safe_load(f) except yaml.YAMLError as exc: raise exc # 1. Setup Data dataset = SignLanguageDataset(config_file["dataset"]["pickle_path"]["train"]) eval_dataset = SignLanguageDataset(config_file["dataset"]["pickle_path"]["dev"]) # 2. Tokenizer tokenizer = AutoTokenizer.from_pretrained( config_file["model"]["pretrained_model_name"] ) # 3. Model Configuration config = SLTConfig( input_dim=config_file["model"].get("input_dim", 832), vocab_size=tokenizer.vocab_size, bos_token_id=tokenizer.cls_token_id, eos_token_id=tokenizer.sep_token_id, pad_token_id=tokenizer.pad_token_id, ) model = SLTModel(config) # 4. Trainer Setup training_args = TrainingArguments(**config_file["training_args"]) collator = SLTDataCollator(tokenizer) trainer = Trainer( model=model, args=training_args, train_dataset=dataset, eval_dataset=eval_dataset, data_collator=collator, ) trainer.train() # Save model and tokenizer model.save_pretrained("./slt_final_model") tokenizer.save_pretrained("./slt_final_model") if __name__ == "__main__": main()