finetune fixes from upstream

This commit is contained in:
Kijai
2024-08-21 17:21:15 +03:00
parent 617c33da0f
commit 9444038bda
2 changed files with 93 additions and 19 deletions
+83 -9
View File
@@ -299,7 +299,10 @@ class FluxTrainer:
training_models = []
params_to_optimize = []
training_models.append(flux)
params_to_optimize.append({"params": list(flux.parameters()), "lr": args.learning_rate})
name_and_params = list(flux.named_parameters())
# single param group for now
params_to_optimize.append({"params": [p for _, p in name_and_params], "lr": args.learning_rate})
param_names = [[n for n, _ in name_and_params]]
# calculate number of trainable parameters
n_params = 0
@@ -455,17 +458,88 @@ class FluxTrainer:
import library.adafactor_fused
library.adafactor_fused.patch_adafactor_fused(optimizer)
for param_group in optimizer.param_groups:
for parameter in param_group["params"]:
double_blocks_to_swap = args.double_blocks_to_swap
single_blocks_to_swap = args.single_blocks_to_swap
num_double_blocks = len(flux.double_blocks)
num_single_blocks = len(flux.single_blocks)
handled_double_block_indices = set()
handled_single_block_indices = set()
for param_group, param_name_group in zip(optimizer.param_groups, param_names):
for parameter, param_name in zip(param_group["params"], param_name_group):
if parameter.requires_grad:
grad_hook = None
def __grad_hook(tensor: torch.Tensor, param_group=param_group):
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
optimizer.step_param(tensor, param_group)
tensor.grad = None
if double_blocks_to_swap:
if param_name.startswith("double_blocks"):
block_idx = int(param_name.split(".")[1])
if (
block_idx not in handled_double_block_indices
and block_idx >= (num_double_blocks - double_blocks_to_swap) - 1
and block_idx < num_double_blocks - 1
):
# swap next (already backpropagated) block
handled_double_block_indices.add(block_idx)
block_idx_cpu = block_idx + 1
block_idx_cuda = double_blocks_to_swap - (num_double_blocks - block_idx_cpu)
parameter.register_post_accumulate_grad_hook(__grad_hook)
# create swap hook
def create_double_swap_grad_hook(bidx, bidx_cuda):
def __grad_hook(tensor: torch.Tensor):
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
optimizer.step_param(tensor, param_group)
tensor.grad = None
# swap blocks if necessary
flux.double_blocks[bidx].to("cpu")
flux.double_blocks[bidx_cuda].to(accelerator.device)
# print(f"Move double block {bidx} to cpu and {bidx_cuda} to device")
return __grad_hook
grad_hook = create_double_swap_grad_hook(block_idx_cpu, block_idx_cuda)
if single_blocks_to_swap:
if param_name.startswith("single_blocks"):
block_idx = int(param_name.split(".")[1])
if (
block_idx not in handled_single_block_indices
and block_idx >= (num_single_blocks - single_blocks_to_swap) - 1
and block_idx < num_single_blocks - 1
):
handled_single_block_indices.add(block_idx)
block_idx_cpu = block_idx + 1
block_idx_cuda = single_blocks_to_swap - (num_single_blocks - block_idx_cpu)
# print(param_name, block_idx_cpu, block_idx_cuda)
# create swap hook
def create_single_swap_grad_hook(bidx, bidx_cuda):
def __grad_hook(tensor: torch.Tensor):
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
optimizer.step_param(tensor, param_group)
tensor.grad = None
# swap blocks if necessary
flux.single_blocks[bidx].to("cpu")
flux.single_blocks[bidx_cuda].to(accelerator.device)
# print(f"Move single block {bidx} to cpu and {bidx_cuda} to device")
return __grad_hook
grad_hook = create_single_swap_grad_hook(block_idx_cpu, block_idx_cuda)
if grad_hook is None:
def __grad_hook(tensor: torch.Tensor, param_group=param_group):
if accelerator.sync_gradients and args.max_grad_norm != 0.0:
accelerator.clip_grad_norm_(tensor, args.max_grad_norm)
optimizer.step_param(tensor, param_group)
tensor.grad = None
grad_hook = __grad_hook
parameter.register_post_accumulate_grad_hook(grad_hook)
elif args.blockwise_fused_optimizers:
# prepare for additional optimizers and lr schedulers
+10 -10
View File
@@ -2,6 +2,9 @@ import math
import torch
from transformers import Adafactor
# stochastic rounding for bfloat16
# The implementation was provided by 2kpr. Thank you very much!
def copy_stochastic_(target: torch.Tensor, source: torch.Tensor):
"""
copies source into target using stochastic rounding
@@ -11,12 +14,7 @@ def copy_stochastic_(target: torch.Tensor, source: torch.Tensor):
source: the target tensor with dtype=float32
"""
# create a random 16 bit integer
result = torch.randint_like(
source,
dtype=torch.int32,
low=0,
high=(1 << 16),
)
result = torch.randint_like(source, dtype=torch.int32, low=0, high=(1 << 16))
# add the random number to the lower 16 bit of the mantissa
result.add_(source.view(dtype=torch.int32))
@@ -29,6 +27,7 @@ def copy_stochastic_(target: torch.Tensor, source: torch.Tensor):
del result
@torch.no_grad()
def adafactor_step_param(self, p, group):
if p.grad is None:
@@ -75,7 +74,7 @@ def adafactor_step_param(self, p, group):
lr = Adafactor._get_lr(group, state)
beta2t = 1.0 - math.pow(state["step"], group["decay_rate"])
update = (grad ** 2) + group["eps"][0]
update = (grad**2) + group["eps"][0]
if factored:
exp_avg_sq_row = state["exp_avg_sq_row"]
exp_avg_sq_col = state["exp_avg_sq_col"]
@@ -105,9 +104,9 @@ def adafactor_step_param(self, p, group):
p_data_fp32.add_(-update)
#if p.dtype in {torch.float16, torch.bfloat16}:
# if p.dtype in {torch.float16, torch.bfloat16}:
# p.copy_(p_data_fp32)
print("param_dtype: ",p.dtype)
if p.dtype == torch.bfloat16:
copy_stochastic_(p, p_data_fp32)
elif p.dtype == torch.float16:
@@ -133,6 +132,7 @@ def adafactor_step(self, closure=None):
return loss
def patch_adafactor_fused(optimizer: Adafactor):
optimizer.step_param = adafactor_step_param.__get__(optimizer)
optimizer.step = adafactor_step.__get__(optimizer)
optimizer.step = adafactor_step.__get__(optimizer)