mirror of
https://github.com/zenlm/enso.git
synced 2026-07-26 22:30:28 +00:00
accumulation debug
This commit is contained in:
+14
-11
@@ -178,11 +178,11 @@ def main(args):
|
||||
pin_memory=True,
|
||||
drop_last=True
|
||||
)
|
||||
logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})")
|
||||
|
||||
model_engine, opt, _, __ = deepspeed.initialize(
|
||||
args=args, model=model, model_parameters=model.parameters())
|
||||
|
||||
logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})\nAccumulation step {model_engine.gradient_accumulation_steps()}")
|
||||
# Variables for monitoring/logging purposes:
|
||||
train_steps = 0
|
||||
log_steps = 0
|
||||
@@ -213,20 +213,23 @@ def main(args):
|
||||
loss = loss_dict["loss"].mean()
|
||||
|
||||
model_engine.backward(loss)
|
||||
model_engine.step()
|
||||
|
||||
if model_engine.is_gradient_accumulation_boundary():
|
||||
model_engine.step()
|
||||
log_steps += 1
|
||||
train_steps += 1
|
||||
|
||||
data_iter_step += 1
|
||||
# Log loss values:
|
||||
running_loss += loss.item()
|
||||
log_steps += 1
|
||||
train_steps += 1
|
||||
if train_steps % args.log_every == 0:
|
||||
|
||||
if train_steps % args.log_every == 0 and train_steps > 0:
|
||||
# Measure training speed:
|
||||
torch.cuda.synchronize()
|
||||
end_time = time()
|
||||
steps_per_sec = log_steps / (end_time - start_time)
|
||||
# Reduce loss history over all processes:
|
||||
avg_loss = torch.tensor(running_loss / log_steps, device=device)
|
||||
avg_loss = torch.tensor(running_loss / data_iter_step, device=device)
|
||||
dist.all_reduce(avg_loss, op=dist.ReduceOp.SUM)
|
||||
avg_loss = avg_loss.item() / dist.get_world_size()
|
||||
logger.info(f"(step={train_steps:07d}) Train Loss: {avg_loss:.4f}, Train Steps/Sec: {steps_per_sec:.2f}")
|
||||
@@ -236,13 +239,13 @@ def main(args):
|
||||
start_time = time()
|
||||
|
||||
# Save DiT checkpoint:
|
||||
if train_steps % args.ckpt_every == 0 and train_steps > 0:
|
||||
if train_steps % args.ckpt_every == 0 and train_steps > 0:
|
||||
try:
|
||||
checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}"
|
||||
model_engine.save_checkpoint(checkpoint_path)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
print(e)
|
||||
|
||||
dist.barrier()
|
||||
|
||||
# model.eval() # important! This disables randomized embedding dropout
|
||||
@@ -261,8 +264,8 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--num-classes", type=int, default=1000)
|
||||
parser.add_argument("--epochs", type=int, default=1400)
|
||||
parser.add_argument("--train_batch_size", type=int, default=2)
|
||||
parser.add_argument("--global-seed", type=int, default=2024)
|
||||
parser.add_argument("--num-workers", type=int, default=0)
|
||||
parser.add_argument("--global-seed", type=int, default=2023)
|
||||
parser.add_argument("--num-workers", type=int, default=4)
|
||||
parser.add_argument("--log-every", type=int, default=100)
|
||||
parser.add_argument('--accum_iter', default=8, type=int,)
|
||||
parser.add_argument('--num_experts', default=8, type=int,)
|
||||
|
||||
Reference in New Issue
Block a user