no_decay = ["bias", "LayerNorm.weight"]
params = [
    {"params": [p for n,p in model.named_parameters() if not any(nd in n for nd in no_decay)], "weight_decay": args.weight_decay},
    {"params": [p for n,p in model.named_parameters() if any(nd in n for nd in no_decay)], "weight_decay": 0.0},
]
optimizer = AdamW(params, lr=args.lr, betas=(0.9, 0.999), eps=1e-8)
total_updates = args.max_steps
warmup_steps = max(1, int(total_updates * args.warmup_ratio))
scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_updates)

scaler = GradScaler(enabled=True)
