File size: 6,118 Bytes
9c5d0ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
#!/usr/bin/env python3
"""
Đếm số tham số (Parameters) trực tiếp từ file model.safetensors của ACT
(Action Chunking Transformer).

Giống count_params_from_file.py (dùng cho Diffusion Policy): chỉ đọc header
JSON (vài KB đầu file) — KHÔNG cần load toàn bộ model vào RAM/VRAM, KHÔNG cần
cài PyTorch.

Thêm so với bản gốc: phân nhóm tham số theo module (backbone thị giác /
attention / feed-forward / embeddings / ...) — đặc thù cho kiến trúc ACT,
để biết phần nào đang chiếm dung lượng khi so sánh các checkpoint (vd. v2 vs v7).

Cách dùng:
    python3 count_params_act.py /duong/dan/toi/model.safetensors
    python3 count_params_act.py model_v2.safetensors model_v7.safetensors   # so sánh nhiều file
"""
import json
import struct
import sys
from math import prod


# Thứ tự nhóm cũng là thứ tự hiển thị trong bảng kết quả.
# Mỗi rule là (tên_nhóm, hàm_kiểm_tra(tên_tensor) -> bool). Rule ở trên được
# ưu tiên trước — vd. "vae_encoder.layers" phải khớp trước "vae_encoder" chung chung.
GROUP_RULES = [
    ("Backbone thị giác (ResNet18)", lambda k: "backbone" in k),
    ("VAE encoder — các tầng Transformer", lambda k: k.startswith("model.vae_encoder.layers")),
    ("VAE encoder — projections/embeddings", lambda k: k.startswith("model.vae_encoder") and "layers" not in k),
    ("Transformer encoder — các tầng", lambda k: k.startswith("model.encoder.layers")),
    ("Transformer encoder — projections/embeddings", lambda k: k.startswith("model.encoder") and "layers" not in k),
    ("Transformer decoder — các tầng", lambda k: k.startswith("model.decoder.layers")),
    ("Transformer decoder — projections/embeddings", lambda k: k.startswith("model.decoder") and "layers" not in k),
    ("Action head (đầu ra hành động)", lambda k: "action_head" in k),
]
GROUP_OTHER = "Khác (chưa phân loại)"


def group_of(tensor_name: str) -> str:
    for group_name, match in GROUP_RULES:
        if match(tensor_name):
            return group_name
    return GROUP_OTHER


def load_header(path: str) -> dict:
    with open(path, "rb") as f:
        # 8 byte đầu = độ dài (uint64, little-endian) của phần header JSON
        header_len = struct.unpack("<Q", f.read(8))[0]
        header = json.loads(f.read(header_len))
    # "__metadata__" không phải tensor, phải loại ra trước khi đếm
    header.pop("__metadata__", None)
    return header


def count_params(path: str) -> dict:
    header = load_header(path)

    total_params = 0
    total_bytes = 0
    dtype_count: dict = {}
    group_params: dict = {}
    group_tensor_count: dict = {}

    for name, info in header.items():
        shape = info["shape"]
        n_elem = prod(shape) if shape else 1  # tensor 0-chiều (scalar) vẫn tính là 1
        total_params += n_elem

        start, end = info["data_offsets"]
        total_bytes += end - start

        dt = info["dtype"]
        dtype_count[dt] = dtype_count.get(dt, 0) + 1

        g = group_of(name)
        group_params[g] = group_params.get(g, 0) + n_elem
        group_tensor_count[g] = group_tensor_count.get(g, 0) + 1

    return {
        "path": path,
        "n_tensors": len(header),
        "dtype_count": dtype_count,
        "total_params": total_params,
        "total_bytes": total_bytes,
        "group_params": group_params,
        "group_tensor_count": group_tensor_count,
    }


def print_report(r: dict) -> None:
    print(f"File:            {r['path']}")
    print(f"Số tensor:       {r['n_tensors']}")
    print(f"Kiểu dữ liệu:    {r['dtype_count']}")
    print(f"Tổng tham số:    {r['total_params']:,}")
    print(f"Kích thước data: {r['total_bytes']:,} byte  ({r['total_bytes']/1024**2:.1f} MiB)")
    print()
    print("--- Phân theo module ---")
    header = f"{'Module':46s}{'#tensor':>9s}{'Tham số':>14s}{'%':>7s}"
    print(header)
    print("-" * len(header))
    ordered = [g for g, _ in GROUP_RULES] + [GROUP_OTHER]
    for g in ordered:
        p = r["group_params"].get(g, 0)
        if p == 0:
            continue
        n = r["group_tensor_count"].get(g, 0)
        pct = 100 * p / r["total_params"] if r["total_params"] else 0
        print(f"{g:46s}{n:9d}{p:14,d}{pct:6.1f}%")
    print()

    # Ước tính VRAM lúc train (FP32): trọng số + gradient + 2 hệ số Adam = 4x.
    # Đây chỉ là phần do THAM SỐ đóng góp — activations (phụ thuộc batch_size
    # và số camera) thường lớn hơn nhiều và KHÔNG được tính ở đây.
    w_mib = r["total_params"] * 4 / 1024**2
    print("--- Ước tính VRAM do THAM SỐ đóng góp (FP32, chưa tính activations) ---")
    print(f"Trọng số:                    {w_mib:8.1f} MiB")
    print(f"Inference (chỉ trọng số):    {w_mib:8.1f} MiB")
    print(f"Training (+grad +2x Adam):   {4*w_mib:8.1f} MiB  (~{4*w_mib/1024:.2f} GiB)")
    print("Lưu ý: activations (ảnh camera) thường lớn hơn khoản này nhiều lần khi train.")


def print_comparison(reports: list) -> None:
    print("=== So sánh nhiều file ===")
    header = f"{'File':30s}{'Tổng tham số':>16s}{'MiB':>10s}"
    print(header)
    print("-" * len(header))
    for r in reports:
        name = r["path"].split("/")[-1]
        print(f"{name:30s}{r['total_params']:16,d}{r['total_bytes']/1024**2:10.1f}")
    if len(reports) == 2:
        a, b = reports
        if b["total_params"]:
            print(f"\nTỉ lệ {a['path'].split('/')[-1]} / {b['path'].split('/')[-1]}: "
                  f"{a['total_params']/b['total_params']:.3f}x")
    print()


if __name__ == "__main__":
    if len(sys.argv) < 2:
        print("Dùng: python3 count_params_act.py <model1.safetensors> [model2.safetensors ...]")
        sys.exit(1)

    all_reports = []
    for p in sys.argv[1:]:
        rpt = count_params(p)
        all_reports.append(rpt)
        print_report(rpt)
        print("=" * 70)
        print()

    if len(all_reports) > 1:
        print_comparison(all_reports)