Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Staged unfreezing and domain-shift audits

Last updated: 7 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

A staged adaptation should show which layer was unfrozen, which population improved and whether the new model forgot behavior it previously handled.

Name the shifted population

Suppose a receipt classifier trained mostly on thermal-paper scans begins receiving glossy mobile-camera images. Keep the original scan validation set and a device-disjoint glossy validation set. Split by receipt and merchant identity where repeated layouts could leak. The question is not whether training loss falls; it is whether the new domain improves without an unacceptable regression on the original one. Identity-safe splitting precedes any adaptation.

Stage the update

First train only the replacement head, then unfreeze a late feature block with a smaller learning rate than the head. Evaluate after each stage and stop if the predeclared acceptance rule fails. Construct optimizer groups only from parameters that require gradients; after a stage transition, rebuild the optimizer deliberately because its groups and moments no longer describe the intended parameter set. If continuity of the head optimizer state matters, migrate it explicitly and verify it rather than silently resetting.

Inspect normalization and preprocessing

Glossy photos may change brightness statistics, but blindly updating BatchNorm running buffers on a tiny domain slice can destabilize older images. Run a controlled variant with those buffers fixed and one with explicit adaptation. Use the same input color space, resize and weight-specific normalization in every variant. Frozen-feature policy explains why gradient flags alone do not define model state.

Use slices that explain failure

Report confusion matrices by paper type, acquisition device and defect class. Track calibration or threshold changes separately from representation changes: a better ranking score can still produce worse decisions at a fixed threshold. If labels are sparse in the shifted domain, include confidence intervals or raw numerator and denominator for each slice. Do not let a broad aggregate hide two missed cutoff receipts in a small critical subgroup.

Record rollback evidence

Save the frozen baseline, each accepted stage and its exact preprocessing and threshold. Re-evaluate both domains after loading from disk to expose buffer or serialization mistakes. Keep the rollout reversible: an inference error or domain regression should select the previous artifact, not trigger another unreviewed training run. Checkpoint state concerns training continuation; serving needs a separately tested inference artifact.

Implementation

python
import torch
from torch import nn
from torchvision.models import resnet18, ResNet18_Weights

domain_model = resnet18(weights=ResNet18_Weights.IMAGENET1K_V1)
domain_model.fc = nn.Linear(domain_model.fc.in_features, 3)
for parameter in domain_model.parameters():
    parameter.requires_grad = False
for parameter in domain_model.fc.parameters():
    parameter.requires_grad = True

def make_optimizer(model: nn.Module, adapt_last_block: bool):
    for parameter in model.layer4.parameters():
        parameter.requires_grad = adapt_last_block
    parameter_groups = [
        {"params": list(model.fc.parameters()), "lr": 0.00043},
    ]
    if adapt_last_block:
        parameter_groups.append({"params": list(model.layer4.parameters()),
                                 "lr": 0.000047})
    optimizer = torch.optim.AdamW(parameter_groups, weight_decay=0.013)
    selected_ids = [id(parameter) for group in optimizer.param_groups
                    for parameter in group["params"]]
    assert len(selected_ids) == len(set(selected_ids))
    return optimizer

frozen_optimizer = make_optimizer(domain_model, adapt_last_block=False)
adapted_optimizer = make_optimizer(domain_model, adapt_last_block=True)
assert len(frozen_optimizer.param_groups) == 1
assert len(adapted_optimizer.param_groups) == 2

Performance and operating cost

Unfreezing the late block adds its backward activations, gradients and optimizer moments. Rebuilding an optimizer at a stage boundary is cheap relative to training, but resetting moments changes the trajectory and must be recorded. A two-domain validation pass roughly doubles evaluation work compared with one domain; it is still cheaper than releasing an adaptation that fails on previously supported receipts.

Common Mistakes

  • Do not infer population improvement from training loss alone.
  • Do not add the same parameter to multiple optimizer groups or leave an unfrozen parameter out.
  • Do not call a stage accepted without checking regression on the original domain.

Read next

Continue the workflow: Project: audit receipt confidence and review handoff.

ai-data
deep-learning
Storage details