#!/usr/bin/env python3 """ Đếm số tham số (Parameters) trực tiếp từ file model.safetensors của Diffusion Policy (LeRobot). 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. An toàn với file nhiều GB. Thêm so với bản gốc (đồng bộ với count_params_act.py): - Phân nhóm tham số theo module (vision encoder / U-Net down-mid-up / FiLM ...) - Ước tính VRAM do THAM SỐ đóng góp (inference và training) - Tách riêng tham số học được và buffer (BatchNorm running stats) — vì safetensors lưu cả hai, nên "tổng phần tử" > "số tham số thực" Cách dùng: python3 count_params_from_file.py /duong/dan/toi/model.safetensors python3 count_params_from_file.py dp_v1.safetensors dp_v2.safetensors # so sánh """ import json import struct import sys from math import prod # Số byte mỗi phần tử theo dtype của safetensors — dùng để ước tính VRAM # cho đúng độ chính xác mà model đang được lưu (FP32 / FP16 / BF16). DTYPE_BYTES = { "F64": 8, "F32": 4, "F16": 2, "BF16": 2, "I64": 8, "I32": 4, "I16": 2, "I8": 1, "U8": 1, "BOOL": 1, } # 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. "rgb_encoder ... backbone" phải khớp trước "rgb_encoder" chung. GROUP_RULES = [ ("Vision encoder — backbone ResNet-18", lambda k: "rgb_encoder" in k and "backbone" in k), ("Vision encoder — SpatialSoftmax + projection", lambda k: "rgb_encoder" in k), ("U-Net — diffusion_step_encoder", lambda k: "diffusion_step_encoder" in k), ("U-Net — down_modules (nhánh nén)", lambda k: "down_modules" in k), ("U-Net — mid_modules (bottleneck)", lambda k: "mid_modules" in k), ("U-Net — up_modules (nhánh giãn)", lambda k: "up_modules" in k), ("U-Net — final_conv (đầu ra action)", lambda k: "final_conv" in k), ("Normalization buffers (mean/std/min/max)", lambda k: "normalize" in k.lower()), ] GROUP_OTHER = "Khác (chưa phân loại)" # Nhóm cắt ngang: FiLM conditioning nằm rải trong mọi residual block của U-Net, # nên không thể tách bằng GROUP_RULES (sẽ trùng với down/mid/up). Thống kê riêng. def is_film(name: str) -> bool: return "cond_encoder" in name def is_buffer(name: str, dtype: str) -> bool: """Buffer = tensor được lưu trong file nhưng KHÔNG phải tham số học được.""" return ( "running_mean" in name or "running_var" in name or "num_batches_tracked" in name or dtype in ("I64", "I32", "BOOL") ) 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_elems = 0 # mọi phần tử được lưu trong file learnable = 0 # tham số học được (đã loại buffer) buffers = 0 # BN running stats, num_batches_tracked, ... total_bytes = 0 film_params = 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 dt = info["dtype"] total_elems += n_elem if is_buffer(name, dt): buffers += n_elem else: learnable += n_elem if is_film(name): film_params += n_elem start, end = info["data_offsets"] total_bytes += end - start 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 # dtype chiếm đa số — dùng để ước tính VRAM đúng với độ chính xác đang lưu dominant_dtype = max(dtype_count, key=dtype_count.get) if dtype_count else "F32" bytes_per_elem = DTYPE_BYTES.get(dominant_dtype, 4) return { "path": path, "n_tensors": len(header), "dtype_count": dtype_count, "dominant_dtype": dominant_dtype, "bytes_per_elem": bytes_per_elem, "total_elems": total_elems, "learnable": learnable, "buffers": buffers, "film_params": film_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 phần tử: {r['total_elems']:,}") print(f" ├─ tham số học được: {r['learnable']:,}") print(f" └─ buffer (BN stats): {r['buffers']:,}") print(f"Kích thước data: {r['total_bytes']:,} byte " f"({r['total_bytes']/1024**2:.1f} MiB / {r['total_bytes']/1e9:.3f} GB)") print() print("--- Phân theo module ---") head = f"{'Module':46s}{'#tensor':>9s}{'Tham số':>14s}{'%':>7s}" print(head) print("-" * len(head)) 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_elems"] if r["total_elems"] else 0 print(f"{g:46s}{n:9d}{p:14,d}{pct:6.1f}%") if r["film_params"]: pct = 100 * r["film_params"] / r["total_elems"] print("-" * len(head)) print(f"{'(cắt ngang) FiLM cond_encoder':46s}" f"{'':>9s}{r['film_params']:14,d}{pct:6.1f}%") print(" ^ nằm rải trong down/mid/up ở trên — KHÔNG cộng thêm vào tổng") print() # Ước tính VRAM lúc train: trọng số + gradient + 2 hệ số Adam (m, v) = 4x. # Đây chỉ là phần do THAM SỐ đóng góp — activations (phụ thuộc batch_size, # số camera, độ phân giải ảnh) thường lớn hơn nhiều và KHÔNG được tính ở đây. b = r["bytes_per_elem"] w_mib = r["total_elems"] * b / 1024**2 print(f"--- Ước tính VRAM do THAM SỐ đóng góp ({r['dominant_dtype']}, chưa tính activations) ---") print(f"Trọng số: {w_mib:8.1f} MiB (~{w_mib/1024:.2f} GiB)") print(f"Inference (chỉ trọng số): {w_mib:8.1f} MiB (~{w_mib/1024:.2f} GiB)") print(f"Training (+grad +2x Adam): {4*w_mib:8.1f} MiB (~{4*w_mib/1024:.2f} GiB)") print() print("Lưu ý khi đọc con số trên:") print(" - Activations (ảnh camera qua ResNet-18) thường lớn hơn khoản này") print(" nhiều lần khi train; đây là sàn dưới, KHÔNG phải tổng VRAM.") print(" - Isaac Sim chạy song song cũng chiếm VRAM đáng kể (render 4 camera).") print(" - Diffusion Policy lặp U-Net num_inference_steps lần mỗi bước inference:") print(" không tăng đỉnh VRAM nhưng gây phân mảnh bộ nhớ -> OOM sớm hơn lý thuyết.") def print_comparison(reports: list) -> None: print("=== So sánh nhiều file ===") head = f"{'File':30s}{'Tổng phần tử':>16s}{'MiB':>10s}" print(head) print("-" * len(head)) for r in reports: name = r["path"].split("/")[-1] print(f"{name:30s}{r['total_elems']:16,d}{r['total_bytes']/1024**2:10.1f}") if len(reports) == 2: a, b = reports if b["total_elems"]: print(f"\nTỉ lệ {a['path'].split('/')[-1]} / {b['path'].split('/')[-1]}: " f"{a['total_elems']/b['total_elems']:.3f}x") print() if __name__ == "__main__": if len(sys.argv) < 2: print("Dùng: python3 count_params_from_file.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)