cleaner merger + use_per_ctx_avg_loss=True as default

This commit is contained in:
51616 2025-09-05 17:52:03 +09:00
parent f02ae874f3
commit 23e502aa65
2 changed files with 18 additions and 17 deletions

View file

@ -428,7 +428,7 @@ class CtxTrainingArguments:
metadata={"help": "Whether to use KL loss."},
)
use_per_ctx_average_loss: bool = field(
default=False,
default=True,
metadata={"help": "Whether to use per-context average loss."},
)
gen_lora_l1_reg_coef: float = field(

View file

@ -21,15 +21,18 @@ def combine_lora(
# Assume all modules share same base rank r
first_module = next(iter(generated_loras))
base_rank = generated_loras[first_module]["A"].shape[-2]
max_rank_needed = compute_rank(n_chunks.max(), base_rank)
sampled_lora = generated_loras[first_module]["A"]
base_rank = sampled_lora.shape[-2]
device = sampled_lora.device
dtype = sampled_lora.dtype
max_rank_needed = int(compute_rank(n_chunks.max(), base_rank))
combined_loras: dict[str, dict[str, Tensor]] = {
module: {"A": None, "B": None} for module in generated_loras.keys()
}
rank_dim = 2
num_groups = len(n_chunks)
rank_per_group = n_chunks * base_rank
rank_per_group = (n_chunks * base_rank).tolist()
for module_name, module_loras in generated_loras.items():
for matrix_key in ("A", "B"):
@ -39,32 +42,30 @@ def combine_lora(
flat_loras = rearrange(
loras, "tot_chunks n_layers r dim -> 1 n_layers (tot_chunks r) dim"
)
per_group_deltas = flat_loras.split(rank_per_group.tolist(), dim=rank_dim)
per_group_deltas = flat_loras.split(rank_per_group, dim=rank_dim)
combined_shape = [num_groups, *per_group_deltas[0].shape[1:]]
combined_shape[rank_dim] = max_rank_needed
combined = torch.zeros(
*combined_shape,
device=per_group_deltas[0].device,
dtype=per_group_deltas[0].dtype,
)
combined = torch.zeros(*combined_shape, device=device, dtype=dtype)
for g, deltas in enumerate(per_group_deltas):
combined_rank = deltas.shape[rank_dim]
# Build slice pattern, slice up to combined_rank.
slice_pattern = [g, slice(None), slice(None), slice(None)]
slice_pattern[rank_dim] = slice(combined_rank)
# slice_pattern = [g, slice(None), slice(None), slice(None)]
# slice_pattern[rank_dim] = slice(combined_rank)
combined[slice_pattern] = deltas
combined[g, :, :combined_rank, :] = deltas
if bias_tensor is not None:
bias_slice_pattern = [g, slice(None), slice(None), slice(None)]
bias_slice_pattern[rank_dim] = slice(
combined_rank, combined_rank + base_rank
# bias_slice_pattern = [g, slice(None), slice(None), slice(None)]
# bias_slice_pattern[rank_dim] = slice(
# combined_rank, combined_rank + base_rank
# )
combined[g, :, combined_rank : combined_rank + base_rank, :] = (
bias_tensor
)
combined[bias_slice_pattern] = bias_tensor
combined_loras[module_name][matrix_key] = combined