Вот код:
Код: Выделить всё
import os
import json
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torch.nn.utils.rnn import pad_sequence
from torch.amp import GradScaler, autocast
# Load tokenized datasets
def load_tokenized_data():
try:
with open('tokenized_en.json', 'r', encoding='utf-8') as ef:
en_tokenized = json.load(ef)
with open('tokenized_ja.json', 'r', encoding='utf-8') as jf:
ja_tokenized = json.load(jf)
if not isinstance(en_tokenized, list) or not isinstance(ja_tokenized, list):
raise ValueError("Both datasets must be lists.")
if len(en_tokenized) != len(ja_tokenized):
raise ValueError("Datasets must have the same length.")
# Filter out empty sentences
filtered_data = [(en, ja) for en, ja in zip(en_tokenized, ja_tokenized) if en and ja]
en_tokenized, ja_tokenized = zip(*filtered_data)
print(f"Filtered out empty sentences. Remaining samples: {len(en_tokenized)}")
return list(en_tokenized), list(ja_tokenized)
except Exception as e:
print(f"Error loading datasets: {e}")
return None, None
# Custom Dataset and DataLoader
class TranslationDataset(Dataset):
def __init__(self, en_sentences, ja_sentences):
self.en_sentences = en_sentences
self.ja_sentences = ja_sentences
def __len__(self):
return len(self.en_sentences)
def __getitem__(self, idx):
en_tensor = torch.tensor(self.en_sentences[idx], dtype=torch.long)
ja_tensor = torch.tensor(self.ja_sentences[idx], dtype=torch.long)
print(f"Data fetched at index {idx}: en_tensor shape {en_tensor.shape}, ja_tensor shape {ja_tensor.shape}")
return en_tensor, ja_tensor
# Updated collate function to ensure same length for src and tgt
def collate_fn(batch):
src_batch, tgt_batch = zip(*batch)
src_batch = pad_sequence(src_batch, batch_first=True, padding_value=0)
tgt_batch = pad_sequence(tgt_batch, batch_first=True, padding_value=0)
# Make sure both src_batch and tgt_batch have the same sequence length
if src_batch.size(1) != tgt_batch.size(1):
max_length = max(src_batch.size(1), tgt_batch.size(1))
src_batch = nn.functional.pad(src_batch, (0, max_length - src_batch.size(1)), value=0)
tgt_batch = nn.functional.pad(tgt_batch, (0, max_length - tgt_batch.size(1)), value=0)
print(f"Collated batch: src_batch shape {src_batch.shape}, tgt_batch shape {tgt_batch.shape}")
return src_batch, tgt_batch
# Function to generate a causal mask
def generate_square_subsequent_mask(sz):
return torch.triu(torch.ones(sz, sz) * float('-inf'), diagonal=1)
# Transformer model with correct attention mask handling
class TransformerModel(nn.Module):
def __init__(self, src_vocab_size, tgt_vocab_size, d_model=512, num_heads=8, num_encoder_layers=6, num_decoder_layers=6):
super(TransformerModel, self).__init__()
self.encoder_embedding = nn.Embedding(src_vocab_size, d_model, padding_idx=0)
self.decoder_embedding = nn.Embedding(tgt_vocab_size, d_model, padding_idx=0)
self.transformer = nn.Transformer(d_model, num_heads, num_encoder_layers, num_decoder_layers, dropout=0.1, batch_first=True)
self.fc_out = nn.Linear(d_model, tgt_vocab_size)
self.num_heads = num_heads
self.d_model = d_model
def generate_attention_mask(self, mask_size, batch_size, device):
# Create the causal mask (for target)
causal_mask = generate_square_subsequent_mask(mask_size).to(torch.float32).to(device)
causal_mask = causal_mask.unsqueeze(0).unsqueeze(0) # Shape (1, 1, mask_size, mask_size)
# Expand for batch_size and num_heads
expanded_mask = causal_mask.expand(batch_size, self.num_heads, mask_size, mask_size)
# Flatten batch_size and num_heads for transformer attention compatibility
final_mask = expanded_mask.reshape(batch_size * self.num_heads, mask_size, mask_size)
return final_mask
def forward(self, src, tgt):
device = src.device
batch_size = src.size(0)
tgt_seq_len = tgt.size(1)
# Embeddings
src_emb = self.encoder_embedding(src)
tgt_emb = self.decoder_embedding(tgt)
# Source and target key padding masks (1 for padding tokens)
src_key_padding_mask = (src == 0).to(device)
tgt_key_padding_mask = (tgt == 0).to(device)
# Causal target mask
tgt_mask = self.generate_attention_mask(tgt_seq_len, batch_size, device)
# Pass through transformer with key padding masks and target mask
transformer_output = self.transformer(
src_emb.transpose(0, 1), # Transformer expects (seq_len, batch, dim)
tgt_emb.transpose(0, 1),
tgt_mask=tgt_mask, # Causal mask for target
src_key_padding_mask=src_key_padding_mask, # Padding mask for src
tgt_key_padding_mask=tgt_key_padding_mask, # Padding mask for tgt
)
# Project output to vocab size
output = self.fc_out(transformer_output.transpose(0, 1))
return output
# Single-GPU setup
def train_model():
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {device}")
# Load datasets
en_tokenized, ja_tokenized = load_tokenized_data()
# Ensure both datasets have sentences
if not en_tokenized or not ja_tokenized:
print("Error: One or both datasets are empty.")
return
print(f"Loaded {len(en_tokenized)} English sentences and {len(ja_tokenized)} Japanese sentences.")
dataset = TranslationDataset(en_tokenized, ja_tokenized)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True, collate_fn=collate_fn, pin_memory=True, num_workers=4)
# Calculate vocab size properly
src_vocab_size = max([max(sentence) for sentence in en_tokenized]) + 1 # +1 to account for 0-based indexing
tgt_vocab_size = max([max(sentence) for sentence in ja_tokenized]) + 1 # +1 to account for 0-based indexing
print(f"Source vocab size: {src_vocab_size}, Target vocab size: {tgt_vocab_size}")
# Initialize the model
model = TransformerModel(src_vocab_size, tgt_vocab_size).to(device)
# Optimizer and Loss function
criterion = nn.CrossEntropyLoss(ignore_index=0) # Ignore padding index
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
# Mixed precision setup
scaler = GradScaler()
# Checkpoint settings
checkpoint_dir = './checkpoints'
best_model_path = os.path.join(checkpoint_dir, 'best_transformer_model.pth')
os.makedirs(checkpoint_dir, exist_ok=True)
best_loss = float('inf')
start_epoch = 0
# Load checkpoint if exists
if os.path.exists(best_model_path):
checkpoint = torch.load(best_model_path, map_location=device)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
start_epoch = checkpoint['epoch'] + 1
best_loss = checkpoint['best_loss']
print(f"Resuming from epoch {start_epoch}")
# Gradient accumulation
accumulation_steps = 4
# Training loop
num_epochs = 10
for epoch in range(start_epoch, num_epochs):
model.train()
total_loss = 0
optimizer.zero_grad()
print(f"Starting epoch {epoch + 1}/{num_epochs}...")
for i, (src, tgt) in enumerate(dataloader):
src, tgt = src.to(device), tgt.to(device)
# Debug: Print batch shapes
print(f"Batch {i}: src shape = {src.shape}, tgt shape = {tgt.shape}")
# Prepare inputs and outputs
tgt_input = tgt[:, :-1] # Input is target without last token
tgt_output = tgt[:, 1:] # Output is target without first token
# Mixed precision: Autocast the forward pass
with autocast(device_type='cuda'):
output = model(src, tgt_input)
loss = criterion(output.view(-1, output.size(-1)), tgt_output.reshape(-1))
# Scale loss and backprop with gradient accumulation
scaler.scale(loss / accumulation_steps).backward()
if (i + 1) % accumulation_steps == 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
total_loss += loss.item()
# Check if any gradients remain to be applied after the final batch
if (i + 1) % accumulation_steps != 0:
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
avg_loss = total_loss / len(dataloader)
print(f"Epoch {epoch + 1}/{num_epochs}, Loss: {avg_loss}")
# Save checkpoint if the model improves
if avg_loss < best_loss:
best_loss = avg_loss
torch.save({
'epoch': epoch,
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'best_loss': best_loss,
}, best_model_path)
print(f"Checkpoint saved: Epoch {epoch + 1}, Loss {avg_loss}")
# Run training
if __name__ == "__main__":
train_model()
Код: Выделить всё
Using device: cuda
Filtered out empty sentences. Remaining samples: 2797388
Loaded 2797388 English sentences and 2797388 Japanese sentences.
Source vocab size: 32000, Target vocab size: 32000
Starting epoch 1/10...
Data fetched at index 1687186: en_tensor shape torch.Size([17]), ja_tensor shape torch.Size([6])Data fetched at index 1061876: en_tensor shape torch.Size([6]), ja_tensor shape torch.Size([5])Data fetched at index 984954: en_tensor shape torch.Size([16]), ja_tensor shape torch.Size([9])Data fetched at index 1700660: en_tensor shape torch.Size([3]), ja_tensor shape torch.Size([6])
Data fetched at index 1922625: en_tensor shape torch.Size([9]), ja_tensor shape torch.Size([6])Data fetched at index 658420: en_tensor shape torch.Size([15]), ja_tensor shape torch.Size([6])Data fetched at index 2708590: en_tensor shape torch.Size([5]), ja_tensor shape torch.Size([4])Data fetched at index 791394: en_tensor shape torch.Size([3]), ja_tensor shape torch.Size([4])
Data fetched at index 1304390: en_tensor shape torch.Size([3]), ja_tensor shape torch.Size([5])Data fetched at index 2495719: en_tensor shape torch.Size([6]), ja_tensor shape torch.Size([5])Data fetched at index 1824387: en_tensor shape torch.Size([5]), ja_tensor shape torch.Size([6])Data fetched at index 2081860: en_tensor shape torch.Size([4]), ja_tensor shape torch.Size([1])
Data fetched at index 304270: en_tensor shape torch.Size([16]), ja_tensor shape torch.Size([11])
Data fetched at index 230481: en_tensor shape torch.Size([9]), ja_tensor shape torch.Size([10])
Data fetched at index 109688: en_tensor shape torch.Size([17]), ja_tensor shape torch.Size([5])Data fetched at index 1955106: en_tensor shape torch.Size([11]), ja_tensor shape torch.Size([10])
Data fetched at index 2602770: en_tensor shape torch.Size([21]), ja_tensor shape torch.Size([9])
Data fetched at index 1239869: en_tensor shape torch.Size([26]), ja_tensor shape torch.Size([18])Data fetched at index 503310: en_tensor shape torch.Size([5]), ja_tensor shape torch.Size([4])
Data fetched at index 2104590: en_tensor shape torch.Size([12]), ja_tensor shape torch.Size([11])
Data fetched at index 2775951: en_tensor shape torch.Size([7]), ja_tensor shape torch.Size([2])Data fetched at index 1320931: en_tensor shape torch.Size([7]), ja_tensor shape torch.Size([5])Data fetched at index 2326649: en_tensor shape torch.Size([13]), ja_tensor shape torch.Size([10])Data fetched at index 1052672: en_tensor shape torch.Size([5]), ja_tensor shape torch.Size([4])
Data fetched at index 2230124: en_tensor shape torch.Size([10]), ja_tensor shape torch.Size([9])Data fetched at index 2019482: en_tensor shape torch.Size([10]), ja_tensor shape torch.Size([3])Data fetched at index 343328: en_tensor shape torch.Size([6]), ja_tensor shape torch.Size([6])Data fetched at index 749484: en_tensor shape torch.Size([12]), ja_tensor shape torch.Size([10])
Data fetched at index 2262498: en_tensor shape torch.Size([7]), ja_tensor shape torch.Size([4])Data fetched at index 2273748: en_tensor shape torch.Size([25]), ja_tensor shape torch.Size([14])Data fetched at index 2433696: en_tensor shape torch.Size([12]), ja_tensor shape torch.Size([6])Data fetched at index 2022148: en_tensor shape torch.Size([10]), ja_tensor shape torch.Size([2])
Data fetched at index 2764622: en_tensor shape torch.Size([6]), ja_tensor shape torch.Size([8])Data fetched at index 757738: en_tensor shape torch.Size([10]), ja_tensor shape torch.Size([8])Data fetched at index 585637: en_tensor shape torch.Size([8]), ja_tensor shape torch.Size([8])Data fetched at index 719355: en_tensor shape torch.Size([8]), ja_tensor shape torch.Size([11])
Data fetch..............
.....Data fetched at index 1162068: en_tensor shape torch.Size([11]), ja_tensor shape torch.Size([11])Data fetched at index 1161735: en_tensor shape torch.Size([7]), ja_tensor shape torch.Size([6])Collated batch: src_batch shape torch.Size([32, 40]), tgt_batch shape torch.Size([32, 40])
Collated batch: src_batch shape torch.Size([32, 27]), tgt_batch shape torch.Size([32, 27])
Data fetched at index 2310066: en_tensor shape torch.Size([7]), ja_tensor shape torch.Size([6])
Collated batch: src_batch shape torch.Size([32, 20]), tgt_batch shape torch.Size([32, 20])
Data fetched at index 2616676: en_tensor shape torch.Size([8]), ja_tensor shape torch.Size([6])
Collated batch: src_batch shape torch.Size([32, 46]), tgt_batch shape torch.Size([32, 46])
Batch 0: src shape = torch.Size([32, 26]), tgt shape = torch.Size([32, 26])
---------------------------------------------------------------------------
RuntimeError Traceback (most recent call last)
Cell In[30], line 229
227 # Run training
228 if __name__ == "__main__":
--> 229 train_model()
Cell In[30], line 194, in train_model()
192 # Mixed precision: Autocast the forward pass
193 with autocast(device_type='cuda'):
--> 194 output = model(src, tgt_input)
195 loss = criterion(output.view(-1, output.size(-1)), tgt_output.reshape(-1))
197 # Scale loss and backprop with gradient accumulation
File /home/zeus/miniconda3/envs/cloudspace/lib/python3.10/site-packages/torch/nn/modules/module.py:1553, in Module._wrapped_call_impl(self, *args, **kwargs)
1551 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc]
1552 else:
-> 1553 return self._call_impl(*args, **kwargs)
File /home/zeus/miniconda3/envs/cloudspace/lib/python3.10/site-packages/torch/nn/modules/module.py:1562, in Module._call_impl(self, *args, **kwargs)
1557 # If we don't have any hooks, we want to skip the rest of the logic in
1558 # this function, and just call forward.
1559 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks
1560 or _global_backward_pre_hooks or _global_backward_hooks
1561 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1562 return forward_call(*args, **kwargs)
1564 try:
1565 result = None
Cell In[30], line 109, in TransformerModel.forward(self, src, tgt)
106 tgt_mask = self.generate_attention_mask(tgt_seq_len, batch_size, device)
108 # Pass through transformer with key padding masks and target mask
--> 109 transformer_output = self.transformer(
110 src_emb.transpose(0, 1), # Transformer expects (seq_len, batch, dim)
111 tgt_emb.transpose(0, 1),
112 tgt_mask=tgt_mask, # Causal mask for target
113 src_key_padding_mask=src_key_padding_mask, # Padding mask for src
114 tgt_key_padding_mask=tgt_key_padding_mask, # Padding mask for tgt
115 )
117 # Project output to vocab size
118 output = self.fc_out(transformer_output.transpose(0, 1))
File /home/zeus/miniconda3/envs/cloudspace/lib/python3.10/site-packages/torch/nn/modules/module.py:1553, in Module._wrapped_call_impl(self, *args, **kwargs)
1551 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc]
1552 else:
-> 1553 return self._call_impl(*args, **kwargs)
File /home/zeus/miniconda3/envs/cloudspace/lib/python3.10/site-packages/torch/nn/modules/module.py:1562, in Module._call_impl(self, *args, **kwargs)
1557 # If we don't have any hooks, we want to skip the rest of the logic in
1558 # this function, and just call forward.
1559 if not (self._backward_hooks or self._backward_pre_hooks or self._forward_hooks or self._forward_pre_hooks
1560 or _global_backward_pre_hooks or _global_backward_hooks
1561 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1562 return forward_call(*args, **kwargs)
1564 try:
1565 result = None
File /home/zeus/miniconda3/envs/cloudspace/lib/python3.10/site-packages/torch/nn/modules/transformer.py:213, in Transformer.forward(self, src, tgt, src_mask, tgt_mask, memory_mask, src_key_padding_mask, tgt_key_padding_mask, memory_key_padding_mask, src_is_causal, tgt_is_causal, memory_is_causal)
211 raise RuntimeError("the batch number of src and tgt must be equal")
212 elif self.batch_first and src.size(0) != tgt.size(0) and is_batched:
--> 213 raise RuntimeError("the batch number of src and tgt must be equal")
215 if src.size(-1) != self.d_model or tgt.size(-1) != self.d_model:
216 raise RuntimeError("the feature number of src and tgt must be equal to d_model")
RuntimeError: the batch number of src and tgt must be equal
Когда я вызываю вывод = модель(src, tgt_input), я хотел, чтобы модель обрабатывала эти парные входные данные без каких-либо ошибок. Моя цель — успешно обучить модель Transformer задаче перевода, поэтому я надеялся на плавный цикл обучения без ошибок времени выполнения, связанных с формами тензоров.
Подробнее здесь: https://stackoverflow.com/questions/790 ... t-be-equal