From 23e502aa65ddabd29ebda852c63fb641f8f7cffe Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 5 Sep 2025 17:52:03 +0900 Subject: [PATCH] cleaner merger + use_per_ctx_avg_loss=True as default --- src/ctx_to_lora/configs.py | 2 +- src/ctx_to_lora/modeling/lora_merger.py | 33 +++++++++++++------------ 2 files changed, 18 insertions(+), 17 deletions(-) diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py index 933f78f..ed3b61b 100644 --- a/src/ctx_to_lora/configs.py +++ b/src/ctx_to_lora/configs.py @@ -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( diff --git a/src/ctx_to_lora/modeling/lora_merger.py b/src/ctx_to_lora/modeling/lora_merger.py index b19b98e..c96cee5 100644 --- a/src/ctx_to_lora/modeling/lora_merger.py +++ b/src/ctx_to_lora/modeling/lora_merger.py @@ -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