File size: 3,052 Bytes
794caa0 1178119 794caa0 1178119 794caa0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 | """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()
|