| |
| from __future__ import annotations |
| import argparse |
| from pathlib import Path |
| import copy |
| import torch |
| import torch.nn.functional as F |
| from tqdm import tqdm |
| from sacflow.utils.config import load_yaml |
| from sacflow.utils.misc import seed_everything, ensure_dir, move_to_device |
| from sacflow.utils.distributed import init_distributed, cleanup, is_main_process |
| from sacflow.data.loader import build_loader |
| from sacflow.models.unet3d import build_model |
| from sacflow.methods.task_space import centered_classifier_basis, project_task_and_residual |
| from sacflow.methods.source_memory import hard_onehot, class_moments, save_source_memory |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--config", required=True) |
| ap.add_argument("--checkpoint", required=True) |
| ap.add_argument("--output", required=True) |
| ap.add_argument("--num-passes", type=int, default=3, help="Number of random-crop passes over source_train for memory estimation") |
| args = ap.parse_args() |
| cfg = load_yaml(args.config) |
| |
| cfg = copy.deepcopy(cfg) |
| cfg.setdefault("data", {}).setdefault("augmentation", {}) |
| cfg["data"]["augmentation"] = {"random_flip": False, "random_intensity_shift": 0.0, "random_intensity_scale": 0.0} |
| seed_everything(int(cfg.get("seed", 1337))) |
| device = init_distributed(cfg.get("distributed", {}).get("backend", "nccl")) |
| model = build_model(cfg).to(device) |
| ckpt = torch.load(args.checkpoint, map_location="cpu") |
| model.load_state_dict(ckpt.get("model", ckpt), strict=False) |
| model.eval() |
| loader = build_loader(cfg, split="source_train", training=True, require_label=True) |
| W = model.final_classifier_weight().detach().to(device) |
| Q = centered_classifier_basis(W) |
| C = cfg["data"]["num_classes"] |
| mu_acc = None |
| var_acc = None |
| feat_mu_acc = None |
| feat_var_acc = None |
| count_acc = None |
| with torch.no_grad(): |
| for pass_idx in range(max(1, args.num_passes)): |
| for batch in tqdm(loader, desc=f"source memory pass {pass_idx+1}/{max(1,args.num_passes)}", disable=not is_main_process()): |
| batch = move_to_device(batch, device) |
| logits, feats = model(batch["image"], return_features=True) |
| feat = feats["prelogit"] |
| task, residual = project_task_and_residual(feat, Q) |
| labels = batch["label"] |
| if labels.shape[-3:] != residual.shape[-3:]: |
| labels = F.interpolate(labels[:,None].float(), size=residual.shape[-3:], mode="nearest")[:,0].long() |
| probs = hard_onehot(labels, C) |
| mu, std, counts = class_moments(residual, probs) |
| fmu, fstd, _ = class_moments(feat, probs) |
| var = std.pow(2) |
| fvar = fstd.pow(2) |
| if mu_acc is None: |
| mu_acc = mu * counts[:,None] |
| var_acc = (var + mu.pow(2)) * counts[:,None] |
| feat_mu_acc = fmu * counts[:,None] |
| feat_var_acc = (fvar + fmu.pow(2)) * counts[:,None] |
| count_acc = counts |
| else: |
| mu_acc += mu * counts[:,None] |
| var_acc += (var + mu.pow(2)) * counts[:,None] |
| feat_mu_acc += fmu * counts[:,None] |
| feat_var_acc += (fvar + fmu.pow(2)) * counts[:,None] |
| count_acc += counts |
| count = count_acc.clamp_min(1e-6) |
| mu = mu_acc / count[:,None] |
| second = var_acc / count[:,None] |
| var = (second - mu.pow(2)).clamp_min(1e-6) |
| fmu = feat_mu_acc / count[:,None] |
| fsecond = feat_var_acc / count[:,None] |
| fvar = (fsecond - fmu.pow(2)).clamp_min(1e-6) |
| mem = { |
| "classifier_weight": W.detach().cpu(), |
| "task_basis_Q": Q.detach().cpu(), |
| "residual_mu": mu.detach().cpu(), |
| "residual_std": torch.sqrt(var).detach().cpu(), |
| "feature_mu": fmu.detach().cpu(), |
| "feature_std": torch.sqrt(fvar).detach().cpu(), |
| "class_counts": count.detach().cpu(), |
| "num_classes": C, |
| "feature_dim": W.shape[1], |
| "note": "Compact source memory. No raw source images stored. Estimated from source_train random crops with augmentation disabled.", |
| "num_passes": args.num_passes, |
| } |
| if is_main_process(): |
| save_source_memory(args.output, mem) |
| print(f"Saved source memory to {args.output}") |
| cleanup() |
|
|
| if __name__ == "__main__": |
| main() |
|
|