"""Check trainable parameters for the DINOv3 ViT + ConvNeXt dual backbone.""" from __future__ import annotations import argparse import json from pathlib import Path import sys ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.append(str(ROOT)) from marine_dual_dinov3_backbone import DINOv3ViTConvNeXtBackbone # noqa: E402 def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--vit-name", default="dinov3_vitl16") parser.add_argument("--convnext-name", default="dinov3_convnext_base") parser.add_argument("--vit-weights", default=None) parser.add_argument("--convnext-weights", default=None) parser.add_argument("--pretrained", action=argparse.BooleanOptionalAction, default=False) parser.add_argument("--output", default=None) return parser.parse_args() def main() -> None: args = parse_args() model = DINOv3ViTConvNeXtBackbone( vit_name=args.vit_name, convnext_name=args.convnext_name, vit_weights=args.vit_weights, convnext_weights=args.convnext_weights, pretrained=args.pretrained, ) trainable = [name for name, param in model.named_parameters() if param.requires_grad] frozen = [name for name, param in model.named_parameters() if not param.requires_grad] summary = model.trainable_summary() report = { "vit_name": args.vit_name, "convnext_name": args.convnext_name, "pretrained": args.pretrained, "summary": summary.__dict__, "trainable_examples": trainable[:80], "frozen_examples": frozen[:40], "checks": { "vit_last_two_blocks_trainable": any(name.startswith("vit.blocks.22.") for name in trainable) and any(name.startswith("vit.blocks.23.") for name in trainable), "vit_earlier_non_norm_blocks_frozen": not any( name.startswith("vit.blocks.21.") and ".norm" not in name for name in trainable ), "vit_early_norm_trainable": any(name.startswith("vit.blocks.0.norm") for name in trainable), "convnext_last_two_stages_trainable": any(name.startswith("convnext.stages.2.") for name in trainable) and any(name.startswith("convnext.stages.3.") for name in trainable), "convnext_early_non_norm_stages_frozen": not any( name.startswith("convnext.stages.0.") and ".norm" not in name for name in trainable ) and not any(name.startswith("convnext.stages.1.") and ".norm" not in name for name in trainable), "convnext_early_norm_trainable": any(name.startswith("convnext.stages.0.") and ".norm" in name for name in trainable), "fusion_trainable": any(name.startswith("fusion.") for name in trainable), }, } text = json.dumps(report, indent=2, ensure_ascii=False) if args.output: Path(args.output).write_text(text, encoding="utf-8") print(text) if __name__ == "__main__": main()