Merge pull request #243 from ElhamDevelopmentStudio/fix/batch-dimension-training

fix: preserve batch dimension in tokenizer and predictor training
This commit is contained in:
ShiYu
2026-04-13 20:38:49 +08:00
committed by GitHub
3 changed files with 14 additions and 14 deletions
+7 -7
View File
@@ -92,9 +92,9 @@ def train_model(model, tokenizer, device, config, save_dir, logger, rank, world_
train_dataset.set_epoch_seed(epoch_idx * 10000 + rank)
valid_dataset.set_epoch_seed(0)
for i, (batch_x, batch_x_stamp) in enumerate(train_loader):
batch_x = batch_x.squeeze(0).to(device, non_blocking=True)
batch_x_stamp = batch_x_stamp.squeeze(0).to(device, non_blocking=True)
for i, (batch_x, batch_x_stamp) in enumerate(train_loader):
batch_x = batch_x.to(device, non_blocking=True)
batch_x_stamp = batch_x_stamp.to(device, non_blocking=True)
# Tokenize input data on-the-fly
with torch.no_grad():
@@ -135,10 +135,10 @@ def train_model(model, tokenizer, device, config, save_dir, logger, rank, world_
model.eval()
tot_val_loss_sum_rank = 0.0
val_batches_processed_rank = 0
with torch.no_grad():
for batch_x, batch_x_stamp in val_loader:
batch_x = batch_x.squeeze(0).to(device, non_blocking=True)
batch_x_stamp = batch_x_stamp.squeeze(0).to(device, non_blocking=True)
with torch.no_grad():
for batch_x, batch_x_stamp in val_loader:
batch_x = batch_x.to(device, non_blocking=True)
batch_x_stamp = batch_x_stamp.to(device, non_blocking=True)
token_seq_0, token_seq_1 = tokenizer.encode(batch_x, half=True)
token_in = [token_seq_0[:, :-1], token_seq_1[:, :-1]]
+5 -5
View File
@@ -123,8 +123,8 @@ def train_model(model, device, config, save_dir, logger, rank, world_size):
train_dataset.set_epoch_seed(epoch_idx * 10000 + rank)
valid_dataset.set_epoch_seed(0) # Keep validation sampling consistent
for i, (ori_batch_x, _) in enumerate(train_loader):
ori_batch_x = ori_batch_x.squeeze(0).to(device, non_blocking=True)
for i, (ori_batch_x, _) in enumerate(train_loader):
ori_batch_x = ori_batch_x.to(device, non_blocking=True)
# --- Gradient Accumulation Loop ---
current_batch_total_loss = 0.0
@@ -174,9 +174,9 @@ def train_model(model, device, config, save_dir, logger, rank, world_size):
model.eval()
tot_val_loss_sum_rank = 0.0
val_sample_count_rank = 0
with torch.no_grad():
for ori_batch_x, _ in val_loader:
ori_batch_x = ori_batch_x.squeeze(0).to(device, non_blocking=True)
with torch.no_grad():
for ori_batch_x, _ in val_loader:
ori_batch_x = ori_batch_x.to(device, non_blocking=True)
zs, _, _, _ = model(ori_batch_x)
_, z = zs
val_loss_item = F.mse_loss(z, ori_batch_x)
+2 -2
View File
@@ -190,7 +190,7 @@ def train_tokenizer(model, device, config, save_dir, logger):
train_sampler.set_epoch(epoch)
for batch_idx, (ori_batch_x, _) in enumerate(train_loader):
ori_batch_x = ori_batch_x.squeeze(0).to(device, non_blocking=True)
ori_batch_x = ori_batch_x.to(device, non_blocking=True)
current_batch_total_loss = 0.0
for j in range(accumulation_steps):
@@ -239,7 +239,7 @@ def train_tokenizer(model, device, config, save_dir, logger):
with torch.no_grad():
for ori_batch_x, _ in val_loader:
ori_batch_x = ori_batch_x.squeeze(0).to(device, non_blocking=True)
ori_batch_x = ori_batch_x.to(device, non_blocking=True)
zs, _, _, _ = (model.module if use_ddp else model)(ori_batch_x)
_, z = zs
val_loss_item = F.mse_loss(z, ori_batch_x)