mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
iterative mode for hypernet
This commit is contained in:
parent
7d95db6672
commit
0e05644d6a
4 changed files with 49 additions and 126 deletions
|
|
@ -99,6 +99,11 @@ if __name__ == "__main__":
|
|||
default=None,
|
||||
help="Number of generated queries to use for context distillation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_iterative_mode",
|
||||
action="store_true",
|
||||
help="Use iterative mode LoRA layer-by-layer generation",
|
||||
)
|
||||
|
||||
cli_args = vars(parser.parse_args())
|
||||
# setup_logging(output_dir, debug=os.getenv("DEBUG", False))
|
||||
|
|
|
|||
|
|
@ -771,6 +771,7 @@ def evaluate(
|
|||
use_flash_attn=True,
|
||||
use_sequence_packing=False, # for generation
|
||||
)
|
||||
model.enable_iterative_mode(args.use_iterative_mode)
|
||||
add_tracker(model.base_model.generate, "generate")
|
||||
add_tracker(model.generate_weights, "generate_weights")
|
||||
add_tracker(model.combine_lora, "combine_lora")
|
||||
|
|
@ -1004,6 +1005,7 @@ def run_eval(
|
|||
cd_update_iterations: int = 10,
|
||||
cd_use_gen_q: bool = False,
|
||||
num_gen_q: int = 20,
|
||||
use_iterative_mode: bool = False,
|
||||
) -> None:
|
||||
"""Run evaluation with the specified parameters."""
|
||||
assert bool(model_name_or_path) ^ bool(checkpoint_path), (
|
||||
|
|
@ -1047,6 +1049,7 @@ def run_eval(
|
|||
# modulated model doesn't see ctx by default
|
||||
# but remove_context has to be false for correct file naming
|
||||
args.remove_context = False
|
||||
args.use_iterative_mode = use_iterative_mode
|
||||
else:
|
||||
args = Namespace(
|
||||
model_name_or_path=model_name_or_path,
|
||||
|
|
|
|||
|
|
@ -98,6 +98,7 @@ class Perceiver(nn.Module):
|
|||
**kwargs,
|
||||
):
|
||||
super().__init__()
|
||||
assert num_extra_modules == 0
|
||||
self.num_layers = num_layers
|
||||
self.num_modules = num_modules
|
||||
self.num_extra_modules = num_extra_modules
|
||||
|
|
@ -133,6 +134,10 @@ class Perceiver(nn.Module):
|
|||
attn_implementation="flash_attention_2",
|
||||
)
|
||||
self.perceiver = Idefics2Perceiver(self.config, self.decoder_config)
|
||||
self.iterative_mode = False
|
||||
|
||||
def enable_iterative_mode(self, x: bool):
|
||||
self.iterative_mode = x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
|
|
@ -141,30 +146,7 @@ class Perceiver(nn.Module):
|
|||
ctx_attn_mask: Integer[Tensor, "bs seq_len"] | None = None,
|
||||
ctx_position_ids: Integer[Tensor, "bs seq_len"] | None = None,
|
||||
):
|
||||
# if len(ctx_features.shape) == 3:
|
||||
# # ["bs seq_len feature_dim"] to
|
||||
# # ["bs n_output_queries feature_dim"]
|
||||
# x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids)
|
||||
# elif len(ctx_features.shape) == 4:
|
||||
# # ["bs n_ctx_model_layers seq_len feature_dim"] to
|
||||
# # to ["bs n_ctx_model_layers n_output_queries feature_dim"]
|
||||
# x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids)
|
||||
# x = rearrange(
|
||||
# x,
|
||||
# (
|
||||
# "bs n_ctx_model_layers (n_modules r) seq_len d -> "
|
||||
# "bs (n_ctx_model_layers n_modules r) d"
|
||||
# ),
|
||||
# n_ctx_model_layers=self.n_ctx_model_layers,
|
||||
# n_modules=self.num_modules,
|
||||
# r=self.r,
|
||||
# )
|
||||
# else:
|
||||
# raise ValueError(
|
||||
# f"ctx_features should be 3D or 4D tensor, got {ctx_features.shape}"
|
||||
# )
|
||||
|
||||
if self.layer_to_layer:
|
||||
if self.layer_to_layer and not self.iterative_mode:
|
||||
if ctx_attn_mask is not None:
|
||||
ctx_attn_mask = repeat(
|
||||
ctx_attn_mask,
|
||||
|
|
@ -186,15 +168,17 @@ class Perceiver(nn.Module):
|
|||
"1 num_layers seq_len feature_dim -> 1 (num_layers seq_len) feature_dim",
|
||||
)
|
||||
|
||||
# print(f"ctx_features shape: {ctx_features.shape}")
|
||||
# print(
|
||||
# f"ctx_attn_mask shape: {ctx_attn_mask.shape if ctx_attn_mask is not None else None}"
|
||||
# )
|
||||
# print(
|
||||
# f"ctx_position_ids shape: {ctx_position_ids.shape if ctx_position_ids is not None else None}"
|
||||
# )
|
||||
x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids)
|
||||
|
||||
if self.layer_to_layer and self.iterative_mode:
|
||||
lora_x = rearrange(
|
||||
x,
|
||||
"bs (n_modules r) d -> bs n_modules r d",
|
||||
n_modules=self.num_modules,
|
||||
r=self.r,
|
||||
)
|
||||
return lora_x, None
|
||||
|
||||
if self.layer_to_layer:
|
||||
per_layer_size = self.num_modules * self.r + self.num_extra_modules
|
||||
x = rearrange(
|
||||
|
|
@ -228,93 +212,10 @@ class Perceiver(nn.Module):
|
|||
n_layers=self.num_layers,
|
||||
)
|
||||
|
||||
# x = rearrange(
|
||||
# x,
|
||||
# "bs (n_layers n_modules) d -> bs n_layers n_modules d",
|
||||
# n_modules=self.num_modules,
|
||||
# n_layers=self.num_layers,
|
||||
# )
|
||||
# lora_emb, extra_emb = unpack(
|
||||
# emb,
|
||||
# [[self.num_modules], [self.num_extra_modules]],
|
||||
# "bs n_layers * feature_dim",
|
||||
# )
|
||||
return lora_x, extra_x
|
||||
|
||||
|
||||
# class Pooler(nn.Module):
|
||||
# def __init__(
|
||||
# self,
|
||||
# feature_size: int,
|
||||
# output_size: int,
|
||||
# pooling_type: POOL_FN,
|
||||
# num_layers: int,
|
||||
# num_modules: int,
|
||||
# *args,
|
||||
# **kwargs,
|
||||
# ):
|
||||
# super().__init__()
|
||||
# self.num_layers = num_layers
|
||||
# self.num_modules = num_modules
|
||||
|
||||
# # NOTE: features will be projected to size = output_size // 2
|
||||
# # then cat with layer and module embeddings (each with size output_size // 4)
|
||||
# # which are collectively form features with size = output_size
|
||||
# self.pool_fn = get_pooling_fn(pooling_type)
|
||||
# self.feature_proj = nn.Linear(feature_size, output_size // 2)
|
||||
# self.ln = nn.LayerNorm(output_size // 2)
|
||||
# self.layer_embs = nn.Sequential(
|
||||
# nn.Embedding(num_layers, output_size // 4),
|
||||
# nn.LayerNorm(output_size // 4),
|
||||
# )
|
||||
# self.module_embs = nn.Sequential(
|
||||
# nn.Embedding(num_modules, output_size // 4),
|
||||
# nn.LayerNorm(output_size // 4),
|
||||
# )
|
||||
# self.mixer = Mixer(output_size, output_size * 4, output_size)
|
||||
# self.mlp = MLPResidualBlock(output_size, output_size * 4, output_size)
|
||||
|
||||
# self.register_buffer("layer_indices", torch.arange(num_layers))
|
||||
# self.register_buffer("module_indices", torch.arange(num_modules))
|
||||
|
||||
# def forward(
|
||||
# self,
|
||||
# features: Float[Tensor, "bs seq_len feature_dim"],
|
||||
# attn_mask: Integer[Tensor, "bs seq_len"] | None = None,
|
||||
# ):
|
||||
# bs = features.shape[0]
|
||||
|
||||
# # [bs, feature_dim]
|
||||
# x = self.ln(self.feature_proj(self.pool_fn(features, attn_mask).float()))
|
||||
# x = repeat(
|
||||
# x,
|
||||
# "bs d -> bs n_layers n_modules d",
|
||||
# n_modules=self.num_modules,
|
||||
# n_layers=self.num_layers,
|
||||
# )
|
||||
|
||||
# layer_embs = self.layer_embs(self.layer_indices) # [num_layers, d]
|
||||
# layer_embs = repeat(
|
||||
# layer_embs,
|
||||
# "n_layers d -> bs n_layers n_modules d",
|
||||
# bs=bs,
|
||||
# n_modules=self.num_modules,
|
||||
# )
|
||||
|
||||
# module_embs = self.module_embs(self.module_indices) # [num_modules, d]
|
||||
# module_embs = repeat(
|
||||
# module_embs,
|
||||
# "n_modules d -> bs n_layers n_modules d",
|
||||
# bs=bs,
|
||||
# n_layers=self.num_layers,
|
||||
# )
|
||||
|
||||
# emb = torch.cat([x, layer_embs, module_embs], dim=3)
|
||||
# return self.mlp(self.mixer(emb))
|
||||
|
||||
|
||||
AGGREGATOR_CLS = {
|
||||
# AGGREGATOR_TYPE.POOLER: Pooler,
|
||||
AGGREGATOR_TYPE.PERCEIVER: Perceiver,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -644,6 +644,10 @@ class HyperLoRA(nn.Module):
|
|||
)
|
||||
return {k: v for k, v in zip(self.extra_modules, layernorms)}
|
||||
|
||||
def enable_iterative_mode(self, x: bool):
|
||||
self.iterative_mode = x
|
||||
self.aggregator.enable_iterative_mode(x)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
features: Float[Tensor, "bs seq_len feature_dim"],
|
||||
|
|
@ -653,15 +657,22 @@ class HyperLoRA(nn.Module):
|
|||
):
|
||||
# [bs, n_layers, n_total_modules, r, feature_dim]
|
||||
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
lora_emb, extra_emb = self.aggregator(features, attn_mask, position_ids)
|
||||
if self.aggregator.layer_to_layer and self.iterative_mode:
|
||||
# iterative inference
|
||||
# features: [bs num_layers seq_len feature_dim]
|
||||
bs, n_layers = features.shape[0:2]
|
||||
lora_emb = torch.empty(
|
||||
(bs, n_layers, self.num_modules, self.r, self.config.latent_size),
|
||||
device=features.device,
|
||||
)
|
||||
for i in range(n_layers):
|
||||
lora_emb[:, i], _ = self.aggregator(
|
||||
features[:, i], attn_mask, position_ids
|
||||
)
|
||||
|
||||
# TODO: add pos emb
|
||||
# TODO:
|
||||
# concat queries from the same chunks
|
||||
# do self-attn
|
||||
# unpack back to [bs, n_layers, n_total_modules, r, feature_dim]
|
||||
# here bs = sum(n_chunks)
|
||||
# we then
|
||||
else:
|
||||
# batched inference
|
||||
lora_emb, _ = self.aggregator(features, attn_mask, position_ids)
|
||||
|
||||
# [bs, n_layers, n_modules, r, max_in_d_outim]
|
||||
flat_loras = None
|
||||
|
|
@ -672,10 +683,10 @@ class HyperLoRA(nn.Module):
|
|||
flat_loras = self.head(norm_lora_emb)
|
||||
|
||||
flat_layernorms = None
|
||||
if self.num_extra_modules:
|
||||
# [bs, n_layers, n_extra_modules, base_hidden_size]
|
||||
extra_emb = self.extra_layers(extra_emb)
|
||||
flat_layernorms = self.extra_head(extra_emb)
|
||||
# if self.num_extra_modules:
|
||||
# # [bs, n_layers, n_extra_modules, base_hidden_size]
|
||||
# extra_emb = self.extra_layers(extra_emb)
|
||||
# flat_layernorms = self.extra_head(extra_emb)
|
||||
|
||||
return flat_loras, flat_layernorms
|
||||
|
||||
|
|
@ -952,6 +963,9 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
ctx_features, ctx_attn_mask, ctx_position_ids
|
||||
)
|
||||
|
||||
def enable_iterative_mode(self, x: bool):
|
||||
self.hypernet.enable_iterative_mode(x)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
ctx_ids: Integer[Tensor, "n_ctx ctx_len"] | None = None,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue