SACFlow / tools /export_source_memory.py
sathiiii's picture
Add tools
93ffd19 verified
Raw
History Blame Contribute Delete
4.56 kB
#!/usr/bin/env python
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)
# Source memory should estimate source feature statistics, not augmentation noise.
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()