#!/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(" 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 [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)