iterative mode for hypernet

This commit is contained in:
51616 2025-09-12 19:05:05 +09:00
parent 7d95db6672
commit 0e05644d6a
4 changed files with 49 additions and 126 deletions

View file

@ -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))

View file

@ -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,

View file

@ -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,
}

View file

@ -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,