| from omegaconf import DictConfig |
| from tqdm import tqdm |
| import torch.nn.functional as F |
|
|
| import clip.clip as clip |
| import torch |
| import torch.nn as nn |
| from torch.utils.data import DataLoader |
|
|
| from .utils import get_class_ids_per_task, get_class_names, batch, merge_we_router, wise_we, moving_avg, l2_loss, \ |
| virtual_vocab |
|
|
| from .cc import conceptual_captions |
|
|
| from . import utils |
| import os |
| import random |
|
|
| from .dynamic_dataset import DynamicDataset |
|
|
|
|
| DFA_R1 = 16 |
| DFA_R2 = 8 |
| DFA_ADAPTER_DROPOUT = 0.1 |
| DFA_ADAPTER_SCALAR = 0.1 |
| GENERIC_TEMPLATES = ( |
| "a photo of a {}", |
| "a clean photo of a {}", |
| "a bad photo of the {}", |
| "a professional photo of a {}", |
| "a studio shot of a {}", |
| "a photo of a {} a realistic setting", |
| "an image of a {}", |
| "a photo of a {} in a photo book", |
| "an image of a {} by a professional photographer", |
| "a photo of a {} using in an experiment", |
| "a photo of a {} by a camera", |
| "a photo of a nice {}", |
| "a photo of the nice {}", |
| "a photo of many {}", |
| "a photo of the cool {}", |
| "a photo of the {}", |
| "art of the {}", |
| "a blurry photo of the {}", |
| "a photo of the {} in a photo book", |
| "a pixelated photo of the {}", |
| ) |
|
|
|
|
| class DFA3Block(nn.Module): |
| """Per-transformer-block DFA-MoE adapter with hierarchical routing. |
| |
| The block follows the paper's functional decoupling: |
| - `e1`: task-agnostic alignment expert (Alignment Pathway) |
| - `e2_list`: task-specific expert group (Plasticity Pathway) |
| - `w_gate_e2` / `w_noise_e2`: inner router over task-specific experts |
| - `router_top`: outer router balancing alignment vs. plasticity outputs |
| |
| Routing is computed from the CLS token and broadcast to the full token sequence. |
| The returned tensor is a residual branch added at the FFN segment. |
| """ |
| def __init__(self, embed_dim: int, e2_top_k: int = 2, num_task_experts: int = 2): |
| super().__init__() |
| self.embed_dim = int(embed_dim) |
| self.e2_top_k = max(1, int(e2_top_k)) |
| self.num_task_experts = max(1, int(num_task_experts)) |
| |
| self.e1 = nn.Sequential( |
| nn.LayerNorm(self.embed_dim), |
| nn.Linear(self.embed_dim, DFA_R1), |
| nn.ReLU(inplace=True), |
| nn.Dropout(p=DFA_ADAPTER_DROPOUT), |
| nn.Linear(DFA_R1, self.embed_dim), |
| ) |
| |
| self.e2_list = nn.ModuleList([ |
| nn.Sequential( |
| nn.LayerNorm(self.embed_dim), |
| nn.Linear(self.embed_dim, DFA_R2), |
| nn.ReLU(inplace=True), |
| nn.Dropout(p=DFA_ADAPTER_DROPOUT), |
| nn.Linear(DFA_R2, self.embed_dim), |
| ) for _ in range(self.num_task_experts) |
| ]) |
| with torch.no_grad(): |
| |
| if isinstance(self.e1[-1], nn.Linear): |
| nn.init.zeros_(self.e1[-1].weight) |
| nn.init.zeros_(self.e1[-1].bias) |
| for mod in self.e2_list: |
| if isinstance(mod[-1], nn.Linear): |
| nn.init.zeros_(mod[-1].weight) |
| nn.init.zeros_(mod[-1].bias) |
| self.router_top = nn.Sequential( |
| nn.LayerNorm(self.embed_dim), |
| nn.Linear(self.embed_dim, 2), |
| ) |
| self.w_gate_e2 = nn.Linear(self.embed_dim, self.num_task_experts, bias=False) |
| self.w_noise_e2 = nn.Linear(self.embed_dim, self.num_task_experts, bias=False) |
|
|
| def _noisy_topk_cls(self, cls_x: torch.Tensor, train: bool = True, noise_epsilon: float = 1e-2) -> torch.Tensor: |
| clean = self.w_gate_e2(cls_x) |
| if train: |
| raw = self.w_noise_e2(cls_x) |
| noise_std = F.softplus(raw) + float(noise_epsilon) |
| logits = clean + torch.randn_like(clean) * noise_std |
| else: |
| logits = clean |
| E = logits.size(-1) |
| k = min(int(self.e2_top_k), int(E)) |
| top_logits, top_idx = logits.topk(k, dim=-1) |
| top_w = torch.softmax(top_logits, dim=-1) |
| if top_w.dtype != logits.dtype: |
| top_w = top_w.to(dtype=logits.dtype) |
| gates = torch.zeros_like(logits, dtype=logits.dtype) |
| gates.scatter_(-1, top_idx, top_w) |
| return gates |
|
|
| def forward(self, pre_ffn: torch.Tensor) -> torch.Tensor: |
| |
| |
| if getattr(self, 'e1_only', False): |
| return DFA_ADAPTER_SCALAR * self.e1(pre_ffn) |
| |
| cls_feat = pre_ffn.permute(1, 0, 2)[:, 0, :] |
| e1_res = DFA_ADAPTER_SCALAR * self.e1(pre_ffn) |
| e2_res_list = [DFA_ADAPTER_SCALAR * m(pre_ffn) for m in self.e2_list] |
| w_top = torch.softmax(self.router_top(cls_feat.float()), dim=-1) |
| w_spec = self._noisy_topk_cls(cls_feat.float(), train=self.training) |
| w_top_b = w_top.unsqueeze(0).unsqueeze(-1) |
| w_spec_b = w_spec.unsqueeze(0).unsqueeze(-1) |
| e2_stack = torch.stack(e2_res_list, dim=2) |
| e2_res = (e2_stack * w_spec_b).sum(dim=2) |
| res_stack = torch.stack([e1_res, e2_res], dim=2) |
| res = (res_stack * w_top_b).sum(dim=2) |
| return res |
|
|
|
|
| class ResidualAttentionBlockWithDFA(nn.Module): |
| """Inject the DFA-MoE residual branch into a CLIP transformer block. |
| |
| The wrapper preserves the original attention residual and augments the FFN |
| residual with the per-block DFA adapter output. |
| """ |
| def __init__(self, base_block: nn.Module, dfa_block: nn.Module): |
| super().__init__() |
| self.base = base_block |
| self.dfa = dfa_block |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| |
| x = x + self.base.attention(self.base.ln_1(x)) |
| |
| pre_ffn = x |
| adapter_out = self.dfa(pre_ffn) |
| x = x + self.base.mlp(self.base.ln_2(x)) + adapter_out |
| return x |
|
|
|
|
| class ClassIncremental(nn.Module): |
| def __init__(self, cfg, device, jit=False): |
| super().__init__() |
| self.prompt_template = cfg.prompt_template |
| self.device = device |
| self.classes_names = None |
| self.model, self.transforms, _ = clip.load(cfg.model_name, device=device, jit=jit) |
| self.ref_model = None |
| |
| if getattr(cfg, 'class_order', None) is not None: |
| self.class_ids_per_task = list(get_class_ids_per_task(cfg)) |
| else: |
| self.class_ids_per_task = None |
| self.current_class_names = [] |
| |
| self.seen_class_ids = [] |
| self.text_tokens = None |
| self.dynamic_dataset = DynamicDataset(cfg) |
| visual_embed_dim = getattr(getattr(self.model, 'visual', None), 'output_dim', None) |
| if visual_embed_dim is None: |
| try: |
| dummy = torch.zeros( |
| 1, |
| 3, |
| getattr(cfg, 'input_resolution', 224), |
| getattr(cfg, 'input_resolution', 224), |
| device=device, |
| ) |
| with torch.no_grad(): |
| visual_embed_dim = self.model.encode_image(dummy).shape[-1] |
| except Exception: |
| visual_embed_dim = getattr(self.model, 'embed_dim', 1024) |
| self.dfa_blocks = [] |
| self._inject_dfa_into_visual_blocks(cfg) |
| |
| self.text_dfa_blocks = [] |
| self._inject_dfa_into_text_blocks(cfg) |
| |
| self.queue_size = int(getattr(cfg, 'moco_queue_size', 4096)) |
| self.register_buffer('queue_img', torch.randn(visual_embed_dim, self.queue_size)) |
| self.register_buffer('queue_txt', torch.randn(visual_embed_dim, self.queue_size)) |
| self.register_buffer('queue_ptr', torch.zeros(1, dtype=torch.long)) |
| self.queue_img = nn.functional.normalize(self.queue_img, dim=0) |
| self.queue_txt = nn.functional.normalize(self.queue_txt, dim=0) |
|
|
| def forward(self, image, _taskid): |
| with torch.no_grad(): |
| logits_per_image, _ = self.model(image, self.text_tokens, 0, is_train=False) |
| probs = logits_per_image.softmax(dim=-1) |
| return probs |
|
|
| def adaptation(self, task_id, cfg, train_dataset, _train_classes_names): |
| |
| task_slice = train_dataset[task_id:task_id + 1] |
| tmp_loader = DataLoader(task_slice, batch_size=256, shuffle=False, num_workers=2) |
| uniq = set() |
| for _inputs, _targets, _tids in tmp_loader: |
| uniq.update(_targets.tolist()) |
| real_ids = sorted(int(x) for x in uniq) |
| |
| start_idx = len(self.current_class_names) |
| |
| if self.class_ids_per_task is None: |
| self.class_ids_per_task = [] |
| while len(self.class_ids_per_task) < task_id: |
| self.class_ids_per_task.append([]) |
| if len(self.class_ids_per_task) == task_id: |
| self.class_ids_per_task.append(real_ids) |
| else: |
| self.class_ids_per_task[task_id] = real_ids |
|
|
| |
| self.current_class_names += get_class_names(self.classes_names, real_ids) |
| self.seen_class_ids += list(real_ids) |
| self.text_tokens = clip.tokenize( |
| [self.prompt_template.format(c) for c in self.current_class_names] |
| ).to(self.device) |
|
|
| |
| self.last_task_real_ids = list(real_ids) |
| self.last_task_start_index = int(start_idx) |
|
|
| if cfg.method != "zeroshot": |
| self.train(task_id, cfg, train_dataset, _train_classes_names) |
|
|
| def train(self, task_id, cfg, train_dataset, _train_classes_names): |
| |
| train_loader = DataLoader( |
| train_dataset[task_id:task_id + 1], |
| batch_size=cfg.batch_size, |
| shuffle=True, |
| num_workers=min(8, os.cpu_count() or 8), |
| pin_memory=True, |
| ) |
|
|
| |
| epochs_global = int(getattr(cfg, 'epochs', 1)) |
| epochs_a = int(getattr(cfg, 'epochs_a', epochs_global)) |
| epochs_b = int(getattr(cfg, 'epochs_b', epochs_global)) |
| total_iterations_a = max(1, epochs_a * len(train_loader)) |
| total_iterations_b = max(1, epochs_b * len(train_loader)) |
| |
| task_class_ids = [int(c) for c in self.class_ids_per_task[task_id]] |
| local_C = len(task_class_ids) |
| max_cid = max(task_class_ids) |
| map_table_cpu = torch.full((max_cid + 1,), -1, dtype=torch.long) |
| for i, cid in enumerate(task_class_ids): |
| map_table_cpu[cid] = i |
|
|
| |
| classnames = get_class_names(self.classes_names, self.class_ids_per_task[task_id]) |
| texts_task = clip.tokenize([self.prompt_template.format(c) for c in classnames]).to(self.device) |
| |
| all_seen_names = [] |
| for tid in range(task_id + 1): |
| all_seen_names.extend(get_class_names(self.classes_names, self.class_ids_per_task[tid])) |
| generic_templates = GENERIC_TEMPLATES |
|
|
| |
| for p in self.model.parameters(): |
| p.requires_grad = False |
|
|
| |
| train_text_e1_stage_a = hasattr(self, 'text_dfa_blocks') and bool(self.text_dfa_blocks) |
| |
| for m in getattr(self, 'dfa_blocks', []): |
| m.e1_only = True |
| for p in m.e1.parameters(): |
| p.requires_grad = True |
| for e2 in getattr(m, 'e2_list', []): |
| for p in e2.parameters(): p.requires_grad = False |
| for p in m.w_gate_e2.parameters(): p.requires_grad = False |
| for p in m.w_noise_e2.parameters(): p.requires_grad = False |
| for p in m.router_top.parameters(): p.requires_grad = False |
| |
| if train_text_e1_stage_a: |
| for m in self.text_dfa_blocks: |
| m.e1_only = True |
| for p in m.e1.parameters(): |
| p.requires_grad = True |
| for e2 in getattr(m, 'e2_list', []): |
| for p in e2.parameters(): p.requires_grad = False |
| for p in m.w_gate_e2.parameters(): p.requires_grad = False |
| for p in m.w_noise_e2.parameters(): p.requires_grad = False |
| for p in m.router_top.parameters(): p.requires_grad = False |
|
|
| lr_a = float(getattr(cfg, 'lr_e1', getattr(cfg, 'lr', 1e-4))) |
| stage_a_base_lrs = [lr_a] |
| |
| img_params_a = [] |
| for m in getattr(self, 'dfa_blocks', []): |
| img_params_a += list(m.e1.parameters()) |
| text_params_a = [] |
| text_lr_a = float(getattr(cfg, 'text_lr_e1', lr_a)) |
| if train_text_e1_stage_a: |
| for m in self.text_dfa_blocks: |
| text_params_a += list(m.e1.parameters()) |
| param_groups_a = [] |
| if img_params_a: |
| param_groups_a.append({ |
| "params": img_params_a, |
| "lr": lr_a, |
| "weight_decay": float(getattr(cfg, 'weight_decay', 0.0)), |
| }) |
| if text_params_a: |
| param_groups_a.append({ |
| "params": text_params_a, |
| "lr": text_lr_a, |
| "weight_decay": float(getattr(cfg, 'weight_decay', 0.0)), |
| }) |
| stage_a_base_lrs.append(text_lr_a) |
| if not param_groups_a: |
| raise RuntimeError("No parameters collected for Stage A optimization in block mode") |
| opt_a = torch.optim.AdamW(param_groups_a) |
| sched_a = utils.cosine_lr(opt_a, stage_a_base_lrs if len(stage_a_base_lrs) > 1 else stage_a_base_lrs[0], 30, total_iterations_a) |
| tau_con = float(getattr(cfg, 'tau_con', 0.05)) |
| use_moco = bool(getattr(cfg, 'use_moco_queue', True)) |
| |
| mq_img_buf: list = [] |
| mq_txt_buf: list = [] |
|
|
| self.model.eval() |
| it = 0 |
| for epoch in range(epochs_a): |
| for inputs, targets_abs, _ in tqdm(train_loader, desc=f"Task {task_id} A (E1)"): |
| sched_a(it) |
| it += 1 |
| inputs = inputs.to(self.device, non_blocking=True) |
| targets_abs = targets_abs.to(self.device, non_blocking=True) |
| |
| map_table = map_table_cpu.to(targets_abs.device) |
| targets = map_table[targets_abs] |
| if (targets < 0).any(): |
| raise RuntimeError("Label mapping failed in Stage A") |
|
|
| |
| templates_cfg = GENERIC_TEMPLATES |
| |
| all_templates = [] |
| base_t = str(getattr(cfg, 'prompt_template', self.prompt_template)) |
| if base_t not in all_templates: |
| all_templates.append(base_t) |
| for t in templates_cfg: |
| if t not in all_templates: |
| all_templates.append(t) |
| Tmpl = len(all_templates) |
|
|
| |
| present_local = torch.unique(targets).tolist() |
| present_local.sort() |
| present_names = [classnames[i] for i in present_local] |
|
|
| texts_present = [tmpl.format(c) for c in present_names for tmpl in all_templates] |
| other_seen = [n for n in dict.fromkeys(all_seen_names) if n not in set(present_names)] |
| texts_neg = [tmpl.format(c) for c in other_seen for tmpl in all_templates] if other_seen else [] |
|
|
| with torch.no_grad(): |
| scale = self.model.logit_scale.exp() |
|
|
| if train_text_e1_stage_a: |
| txt_present_anchors = self.model.encode_text(clip.tokenize(texts_present).to(self.device)) |
| txt_present_anchors = txt_present_anchors / txt_present_anchors.norm(dim=-1, keepdim=True) |
| if texts_neg: |
| txt_neg = self.model.encode_text(clip.tokenize(texts_neg).to(self.device)) |
| txt_neg = txt_neg / txt_neg.norm(dim=-1, keepdim=True) |
| else: |
| txt_neg = None |
| else: |
| with torch.no_grad(): |
| txt_present_anchors = self.model.encode_text(clip.tokenize(texts_present).to(self.device)) |
| txt_present_anchors = txt_present_anchors / txt_present_anchors.norm(dim=-1, keepdim=True) |
| if texts_neg: |
| txt_neg = self.model.encode_text(clip.tokenize(texts_neg).to(self.device)) |
| txt_neg = txt_neg / txt_neg.norm(dim=-1, keepdim=True) |
| else: |
| txt_neg = None |
|
|
| |
| img0_grad = self.model.encode_image(inputs) |
| feats = img0_grad |
| feats = feats / feats.norm(dim=-1, keepdim=True) |
| |
| present_index_map = {int(lid): idx for idx, lid in enumerate(present_local)} |
|
|
| |
| B = feats.size(0) |
| Tmpl = max(1, Tmpl) |
| txt_pos_mean = [] |
| for i in range(B): |
| cls_lid = int(targets[i].item()) |
| pidx = present_index_map[cls_lid] |
| start = pidx * Tmpl |
| end = start + Tmpl |
| txt_block = txt_present_anchors[start:end] |
| txt_pos_mean.append(txt_block.mean(dim=0, keepdim=True)) |
| txt_pos_mean = torch.cat(txt_pos_mean, dim=0) |
|
|
| opt_a.zero_grad(set_to_none=True) |
| |
| qcnt_a = int(getattr(self, 'queue_count', 0)) |
| if use_moco and self.queue_size > 0 and qcnt_a > 0: |
| k_valid_a = min(qcnt_a, int(self.queue_size)) |
| queue_txt_batch = self.queue_txt[:, :k_valid_a].clone().detach().t().to(feats.device) |
| cand_txt = torch.cat( |
| [ |
| txt_present_anchors, |
| txt_neg if txt_neg is not None else txt_present_anchors.new_zeros((0, txt_present_anchors.size(1))), |
| queue_txt_batch, |
| ], |
| dim=0, |
| ) |
| queue_img_batch = self.queue_img[:, :k_valid_a].clone().detach().t().to(feats.device) |
| cand_img = torch.cat([feats, queue_img_batch], dim=0) |
| else: |
| cand_txt = torch.cat( |
| [ |
| txt_present_anchors, |
| txt_neg if txt_neg is not None else txt_present_anchors.new_zeros((0, txt_present_anchors.size(1))), |
| ], |
| dim=0, |
| ) |
| cand_img = feats |
|
|
| |
| N_txt = cand_txt.size(0) |
| pos_mask_i2t = torch.zeros(B, N_txt, dtype=torch.bool, device=feats.device) |
| for i in range(B): |
| cls_lid = int(targets[i].item()) |
| pidx = present_index_map[cls_lid] |
| start = pidx * Tmpl |
| end = start + Tmpl |
| pos_mask_i2t[i, start:end] = True |
|
|
| logits_i2t = (scale * feats @ cand_txt.t()) / max(1e-6, tau_con) |
| |
| pos_logits = logits_i2t.masked_fill(~pos_mask_i2t, float('-inf')) |
| numer = torch.logsumexp(pos_logits, dim=1) |
| denom = torch.logsumexp(logits_i2t, dim=1) |
| loss_i2t = -(numer - denom).mean() |
|
|
| |
| A = txt_present_anchors.size(0) |
| N_img = cand_img.size(0) |
| pos_mask_t2i = torch.zeros(A, N_img, dtype=torch.bool, device=feats.device) |
| |
| for a in range(A): |
| cls_local = present_local[a // Tmpl] |
| |
| match_idx = (targets == int(cls_local)).nonzero(as_tuple=False).squeeze(1) |
| if match_idx.numel() > 0: |
| pos_mask_t2i[a, match_idx] = True |
|
|
| logits_t2i = (scale * txt_present_anchors @ cand_img.t()) / max(1e-6, tau_con) |
| pos_logits_t2i = logits_t2i.masked_fill(~pos_mask_t2i, float('-inf')) |
| numer_t2i = torch.logsumexp(pos_logits_t2i, dim=1) |
| denom_t2i = torch.logsumexp(logits_t2i, dim=1) |
| |
| valid_rows = pos_mask_t2i.any(dim=1) |
| if valid_rows.any(): |
| loss_t2i = -((numer_t2i[valid_rows] - denom_t2i[valid_rows]).mean()) |
| else: |
| loss_t2i = torch.zeros((), device=feats.device, dtype=feats.dtype) |
|
|
| loss_a = (loss_i2t + loss_t2i) * 0.5 |
| loss_a.backward() |
| opt_a.step() |
| |
| if use_moco and self.queue_size > 0: |
| try: |
| mq_img_buf.append(feats.detach().cpu()) |
| mq_txt_buf.append(txt_pos_mean.detach().cpu()) |
| except Exception: |
| pass |
| |
| for m in getattr(self, 'dfa_blocks', []): |
| m.e1_only = False |
| if train_text_e1_stage_a: |
| for m in self.text_dfa_blocks: |
| m.e1_only = False |
| |
|
|
| |
| |
| lr_e2 = float(getattr(cfg, 'lr_e2', getattr(cfg, 'lr', 1e-3))) |
| lr_e2_text = float(getattr(cfg, 'text_lr_e2', max(lr_e2 * 0.1, 1e-6))) |
| lr_e2_router = float(getattr(cfg, 'lr_e2_router', lr_e2)) |
| lr_top_router = float(getattr(cfg, 'lr_top_router', 5.0e-6)) |
| weight_decay = float(getattr(cfg, 'weight_decay', 0.0)) |
|
|
| |
| img_adapt_params, txt_adapt_params = [], [] |
| e2_router_params = [] |
| for m in self.dfa_blocks: |
| if hasattr(m, 'e1'): |
| for p in m.e1.parameters(): p.requires_grad = False |
| for e2 in getattr(m, 'e2_list', []): |
| for p in e2.parameters(): p.requires_grad = True |
| img_adapt_params += list(e2.parameters()) |
| for p in m.w_gate_e2.parameters(): p.requires_grad = True |
| for p in m.w_noise_e2.parameters(): p.requires_grad = True |
| e2_router_params += list(m.w_gate_e2.parameters()) + list(m.w_noise_e2.parameters()) |
| if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks: |
| for m in self.text_dfa_blocks: |
| if hasattr(m, 'e1'): |
| for p in m.e1.parameters(): p.requires_grad = False |
| for e2 in getattr(m, 'e2_list', []): |
| for p in e2.parameters(): p.requires_grad = True |
| txt_adapt_params += list(e2.parameters()) |
| for p in m.w_gate_e2.parameters(): p.requires_grad = True |
| for p in m.w_noise_e2.parameters(): p.requires_grad = True |
| e2_router_params += list(m.w_gate_e2.parameters()) + list(m.w_noise_e2.parameters()) |
|
|
| |
| top_router_params = [] |
| if hasattr(self, 'dfa_blocks') and self.dfa_blocks: |
| for m in self.dfa_blocks: |
| for p in m.router_top.parameters(): p.requires_grad = True |
| top_router_params += list(m.router_top.parameters()) |
| if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks: |
| for m in self.text_dfa_blocks: |
| for p in m.router_top.parameters(): p.requires_grad = True |
| top_router_params += list(m.router_top.parameters()) |
|
|
| param_groups = [] |
| base_lrs = [] |
| if img_adapt_params: |
| param_groups.append({"params": img_adapt_params, "lr": lr_e2, "weight_decay": weight_decay}) |
| base_lrs.append(lr_e2) |
| if txt_adapt_params: |
| param_groups.append({"params": txt_adapt_params, "lr": lr_e2_text, "weight_decay": weight_decay}) |
| base_lrs.append(lr_e2_text) |
| if e2_router_params: |
| param_groups.append({"params": e2_router_params, "lr": lr_e2_router, "weight_decay": weight_decay}) |
| base_lrs.append(lr_e2_router) |
| if top_router_params: |
| param_groups.append({"params": top_router_params, "lr": lr_top_router, "weight_decay": weight_decay}) |
| base_lrs.append(lr_top_router) |
|
|
| if not param_groups: |
| raise RuntimeError("No parameters collected for Stage B optimization in block mode") |
| opt_b = torch.optim.AdamW(param_groups) |
| sched_b = utils.cosine_lr(opt_b, base_lrs if len(base_lrs) > 1 else base_lrs[0], 30, total_iterations_b) |
|
|
| it = 0 |
| for epoch in range(epochs_b): |
| for inputs, targets_abs, _ in tqdm(train_loader, desc=f"Task {task_id} B (E2+Router)"): |
| sched_b(it) |
| it += 1 |
| inputs = inputs.to(self.device, non_blocking=True) |
| targets_abs = targets_abs.to(self.device, non_blocking=True) |
| map_table = map_table_cpu.to(targets_abs.device) |
| targets = map_table[targets_abs] |
| if (targets < 0).any(): |
| raise RuntimeError("Label mapping failed in Stage B") |
|
|
| |
| if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks: |
| txt = self.model.encode_text(texts_task) |
| txt = txt / txt.norm(dim=-1, keepdim=True) |
| with torch.no_grad(): |
| scale = self.model.logit_scale.exp() |
| else: |
| with torch.no_grad(): |
| txt = self.model.encode_text(texts_task) |
| txt = txt / txt.norm(dim=-1, keepdim=True) |
| scale = self.model.logit_scale.exp() |
|
|
| opt_b.zero_grad(set_to_none=True) |
| img = self.model.encode_image(inputs) |
| img = img / img.norm(dim=-1, keepdim=True) |
| fused = img |
| fused = fused / fused.norm(dim=-1, keepdim=True) |
| logits = scale * fused @ txt.t() |
| loss_ce = F.cross_entropy(logits[:, :local_C], targets, label_smoothing=float(getattr(cfg, 'ls', 0.0))) |
| loss = loss_ce |
| |
| lambda_b_con = float(getattr(cfg, 'lambda_b_con', 0.0)) |
| if lambda_b_con > 0.0 and local_C > 1 and hasattr(self, 'queue_img') and hasattr(self, 'queue_txt'): |
| tau_b_con = float(getattr(cfg, 'tau_b_con', getattr(cfg, 'tau_con', 0.05))) |
| qcnt_con = int(getattr(self, 'queue_count', 0)) |
| B_con = fused.size(0) |
| |
| if qcnt_con > 0 and self.queue_size > 0: |
| k_valid_con = min(qcnt_con, int(self.queue_size)) |
| queue_txt_neg = self.queue_txt[:, :k_valid_con].t().to(fused.device) |
| queue_txt_neg = queue_txt_neg / (queue_txt_neg.norm(dim=-1, keepdim=True) + 1e-12) |
| cand_txt_b = torch.cat([txt, queue_txt_neg], dim=0) |
| queue_img_neg = self.queue_img[:, :k_valid_con].t().to(fused.device) |
| queue_img_neg = queue_img_neg / (queue_img_neg.norm(dim=-1, keepdim=True) + 1e-12) |
| cand_img_b = torch.cat([fused, queue_img_neg], dim=0) |
| else: |
| cand_txt_b = txt |
| cand_img_b = fused |
| |
| N_txt_b = cand_txt_b.size(0) |
| logits_i2t_b = (scale * fused @ cand_txt_b.t()) / max(1e-6, tau_b_con) |
| |
| pos_mask_i2t_b = torch.zeros(B_con, N_txt_b, dtype=torch.bool, device=fused.device) |
| for i in range(B_con): |
| pos_idx = int(targets[i].item()) |
| if 0 <= pos_idx < local_C: |
| pos_mask_i2t_b[i, pos_idx] = True |
| pos_logits_i2t_b = logits_i2t_b.masked_fill(~pos_mask_i2t_b, float('-inf')) |
| numer_i2t_b = torch.logsumexp(pos_logits_i2t_b, dim=1) |
| denom_i2t_b = torch.logsumexp(logits_i2t_b, dim=1) |
| loss_i2t_b = -(numer_i2t_b - denom_i2t_b).mean() |
| |
| N_img_b = cand_img_b.size(0) |
| logits_t2i_b = (scale * txt @ cand_img_b.t()) / max(1e-6, tau_b_con) |
| pos_mask_t2i_b = torch.zeros(local_C, N_img_b, dtype=torch.bool, device=fused.device) |
| for c in range(local_C): |
| match_idx = (targets == c).nonzero(as_tuple=False).squeeze(1) |
| if match_idx.numel() > 0: |
| pos_mask_t2i_b[c, match_idx] = True |
| pos_logits_t2i_b = logits_t2i_b.masked_fill(~pos_mask_t2i_b, float('-inf')) |
| numer_t2i_b = torch.logsumexp(pos_logits_t2i_b, dim=1) |
| denom_t2i_b = torch.logsumexp(logits_t2i_b, dim=1) |
| valid_rows_b = pos_mask_t2i_b.any(dim=1) |
| if valid_rows_b.any(): |
| loss_t2i_b = -((numer_t2i_b[valid_rows_b] - denom_t2i_b[valid_rows_b]).mean()) |
| else: |
| loss_t2i_b = torch.zeros((), device=fused.device, dtype=fused.dtype) |
| loss_b_con = (loss_i2t_b + loss_t2i_b) * 0.5 |
| loss = loss + lambda_b_con * loss_b_con |
| loss.backward() |
| opt_b.step() |
|
|
| |
| if use_moco and self.queue_size > 0: |
| if 'mq_img_buf' in locals() and len(mq_img_buf) > 0: |
| try: |
| img_cat = torch.cat(mq_img_buf, dim=0) |
| txt_cat = torch.cat(mq_txt_buf, dim=0) |
| |
| if int(task_id) >= 1: |
| task_quota = int(getattr(cfg, 'moco_task_enqueue_quota', 32)) |
| task_quota = max(0, min(task_quota, int(self.queue_size))) |
| if task_quota > 0: |
| if img_cat.size(0) > task_quota: |
| sample_idx = torch.randperm(img_cat.size(0), device=img_cat.device)[:task_quota] |
| img_cat = img_cat[sample_idx] |
| txt_cat = txt_cat[sample_idx] |
| if img_cat.numel() > 0: |
| img_cat = img_cat.to(self.device, non_blocking=True) |
| txt_cat = txt_cat.to(self.device, non_blocking=True) |
| |
| chunk = min(1024, self.queue_size) |
| for s in range(0, img_cat.size(0), chunk): |
| self._dequeue_and_enqueue(img_cat[s:s+chunk], txt_cat[s:s+chunk]) |
| except Exception: |
| pass |
|
|
| def compute_logits(self, inputs, text_tokens): |
| """Run the unified DFA-MoE inference path against the provided text bank.""" |
| m_prev = self.model.training |
| b_prev = [m.training for m in getattr(self, 'dfa_blocks', [])] |
| t_prev = [m.training for m in getattr(self, 'text_dfa_blocks', [])] if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks else [] |
| self.model.eval() |
| for m in getattr(self, 'dfa_blocks', []): |
| m.eval() |
| if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks: |
| for m in self.text_dfa_blocks: |
| m.eval() |
| with torch.no_grad(): |
| img = self.model.encode_image(inputs) |
| img = img / img.norm(dim=-1, keepdim=True) |
| |
| fused = img |
| fused = fused / fused.norm(dim=-1, keepdim=True) |
| txt = self.model.encode_text(text_tokens) |
| txt = txt / txt.norm(dim=-1, keepdim=True) |
| scale = self.model.logit_scale.exp() |
| logits = scale * fused @ txt.t() |
| if hasattr(self, 'text_dfa_blocks') and self.text_dfa_blocks: |
| for m, s in zip(self.text_dfa_blocks, t_prev): |
| m.train(s) |
| for m, s in zip(getattr(self, 'dfa_blocks', []), b_prev): |
| m.train(s) |
| self.model.train(m_prev) |
| return logits |
|
|
| @torch.no_grad() |
| def _dequeue_and_enqueue(self, img_feats: torch.Tensor, txt_feats: torch.Tensor): |
| """Update the historical FIFO feature queue used by both training stages.""" |
| if not hasattr(self, 'queue_count'): |
| try: |
| self.queue_count = 0 |
| except Exception: |
| self.queue_count = 0 |
| batch_size = img_feats.size(0) |
| ptr = int(self.queue_ptr) |
| |
| if ptr + batch_size <= self.queue_size: |
| self.queue_img[:, ptr:ptr + batch_size] = img_feats.t() |
| self.queue_txt[:, ptr:ptr + batch_size] = txt_feats.t() |
| ptr = (ptr + batch_size) % self.queue_size |
| else: |
| |
| remain = self.queue_size - ptr |
| self.queue_img[:, ptr:] = img_feats[:remain].t() |
| self.queue_txt[:, ptr:] = txt_feats[:remain].t() |
| overflow = batch_size - remain |
| if overflow > 0: |
| self.queue_img[:, :overflow] = img_feats[remain:].t() |
| self.queue_txt[:, :overflow] = txt_feats[remain:].t() |
| ptr = overflow |
| self.queue_ptr[0] = ptr |
| |
| try: |
| self.queue_count = min(self.queue_size, int(self.queue_count) + batch_size) |
| except Exception: |
| self.queue_count = self.queue_size |
|
|
| def _inject_dfa_into_visual_blocks(self, cfg): |
| """Inject per-block DFA-MoE adapters into the visual transformer.""" |
| visual = getattr(self.model, 'visual', None) |
| transformer = getattr(visual, 'transformer', None) |
| resblocks = getattr(transformer, 'resblocks', None) |
| if visual is None or transformer is None or resblocks is None: |
| return |
| |
| new_blocks = [] |
| self.dfa_blocks = [] |
| top_k = int(getattr(cfg, 'e2_top_k', 2)) |
| for b in resblocks: |
| |
| width = None |
| try: |
| width = int(getattr(b, 'ln_1').weight.shape[0]) |
| except Exception: |
| width = int(getattr(visual, 'width', 768)) |
| dfab = DFA3Block( |
| width, |
| e2_top_k=top_k, |
| num_task_experts=int(getattr(cfg, 'num_task_experts', 2)), |
| ).to(self.device) |
| wrap = ResidualAttentionBlockWithDFA(b, dfab) |
| new_blocks.append(wrap) |
| self.dfa_blocks.append(dfab) |
| |
| transformer.resblocks = nn.Sequential(*new_blocks) |
|
|
| def _inject_dfa_into_text_blocks(self, cfg): |
| """Inject the same DFA-MoE structure into the text transformer.""" |
| text_tf = getattr(self.model, 'transformer', None) |
| resblocks = getattr(text_tf, 'resblocks', None) |
| if text_tf is None or resblocks is None: |
| return |
| |
| top_k = int(getattr(cfg, 'text_e2_top_k', getattr(cfg, 'e2_top_k', 2))) |
|
|
| |
| new_blocks = [] |
| self.text_dfa_blocks = [] |
| blocks = list(resblocks) |
| for b in blocks: |
| |
| try: |
| width = int(getattr(b, 'ln_1').weight.shape[0]) |
| except Exception: |
| |
| width = int(getattr(text_tf, 'width', 512)) |
| dfab = DFA3Block( |
| width, |
| e2_top_k=top_k, |
| num_task_experts=int(getattr(cfg, 'num_task_experts', 2)), |
| ).to(self.device) |
| wrap = ResidualAttentionBlockWithDFA(b, dfab) |
| new_blocks.append(wrap) |
| self.text_dfa_blocks.append(dfab) |
| text_tf.resblocks = nn.Sequential(*new_blocks) |
|
|
| class DomainIncremental(nn.Module): |
| pass |
|
|
|
|
| class TaskAgnostic(nn.Module): |
| pass |
|
|
|
|
| def load_model(cfg: DictConfig, device: torch.device) -> nn.Module: |
| r"""Instantiate the continual-learning model for the requested scenario.""" |
| if cfg.scenario == "class": |
| return ClassIncremental(cfg, device) |
| elif cfg.scenario == "domain": |
| return DomainIncremental(cfg, device) |
| elif cfg.scenario == "task-aganostic": |
| return TaskAgnostic(cfg, device) |
| else: |
| raise ValueError(f""" |
| `{cfg.scenarios}` is not a valid scenario, |
| Please choose from ['class', "domain', 'task-agnostic'] |
| """) |
|
|