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