| """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 |
|
|
|
|
| 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() |
|
|