mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
cleaner merger + use_per_ctx_avg_loss=True as default
This commit is contained in:
parent
f02ae874f3
commit
23e502aa65
2 changed files with 18 additions and 17 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue