EADST

Print Transformers Pytorch Model Information

import os
import re
import torch
from safetensors import safe_open
from safetensors.torch import load_file
import glob
from collections import defaultdict
import numpy as np

model_dir = "/dfs/data/model_path_folder/"

def inspect_model_weights(directory_path):
    """
    检索文件夹中所有的bin或safetensors文件并打印模型权重信息

    参数:
        directory_path (str): 包含模型文件的文件夹路径
    """
    # 查找所有bin和safetensors文件
    bin_files = glob.glob(os.path.join(directory_path, "*.bin"))
    safetensors_files = glob.glob(os.path.join(directory_path, "*.safetensors"))

    all_files = bin_files + safetensors_files

    if not all_files:
        print(f"在 {directory_path} 中没有找到bin或safetensors文件")
        return

    print(f"找到 {len(all_files)} 个模型文件:")
    for idx, file_path in enumerate(all_files):
        print(f"{idx+1}. {os.path.basename(file_path)}")

    total_size = 0
    param_count = 0
    layer_stats = defaultdict(int)
    tensor_types = defaultdict(int)
    shape_info = defaultdict(list)

    # 处理每个文件
    for file_path in all_files:
        file_size = os.path.getsize(file_path) / (1024 * 1024)  # MB
        total_size += file_size

        print(f"\n检查文件: {os.path.basename(file_path)} ({file_size:.2f} MB)")

        # 根据文件扩展名加载权重
        if file_path.endswith('.bin'):
            try:
                weights = torch.load(file_path, map_location='cpu')
            except Exception as e:
                print(f"  无法加载 {file_path}: {e}")
                continue
        else:  # safetensors
            try:
                weights = load_file(file_path)
            except Exception as e:
                print(f"  无法加载 {file_path}: {e}")
                continue

        # 分析权重
        print(f"  包含 {len(weights)} 个张量")
        for key, tensor in weights.items():
            # 统计参数数量
            num_params = np.prod(tensor.shape)
            param_count += num_params

            # 统计层类型
            layer_type = "other"
            if "attention" in key or "attn" in key:
                layer_type = "attention"
            elif "mlp" in key or "ffn" in key:
                layer_type = "feed_forward"
            elif "embed" in key:
                layer_type = "embedding"
            elif "norm" in key or "ln" in key:
                layer_type = "normalization"
            layer_stats[layer_type] += num_params

            # 统计张量类型
            tensor_types[tensor.dtype] += num_params

            # 记录形状信息
            shape_str = str(tensor.shape)
            shape_info[shape_str].append(key)

            # 打印详细信息(前10个张量)
            if len(shape_info) <= 10 or num_params > 1_000_000:
                print(f"  - {key}: 形状={tensor.shape}, 类型={tensor.dtype}, 参数数={num_params:,}")

    # 打印汇总信息
    print("\n模型权重汇总:")
    print(f"总文件大小: {total_size:.2f} MB")
    print(f"总参数数量: {param_count:,}")

    print("\n按层类型划分的参数:")
    for layer_type, count in layer_stats.items():
        percentage = (count / param_count) * 100
        print(f"  {layer_type}: {count:,} 参数 ({percentage:.2f}%)")

    print("\n张量数据类型分布:")
    for dtype, count in tensor_types.items():
        percentage = (count / param_count) * 100
        print(f"  {dtype}: {count:,} 参数 ({percentage:.2f}%)")

    print("\n常见张量形状:")
    sorted_shapes = sorted(shape_info.items(), key=lambda x: np.prod(eval(x[0])), reverse=True)
    for i, (shape, keys) in enumerate(sorted_shapes[:10]):
        num_params = np.prod(eval(shape))
        percentage = (num_params * len(keys) / param_count) * 100
        print(f"  {shape}: {len(keys)} 个张量, 每个 {num_params:,} 参数 (总共占 {percentage:.2f}%)")
        if i < 3:  # 只显示前3种最常见形状的示例
            print(f"    例如: {', '.join(keys[:3])}" + ("..." if len(keys) > 3 else ""))

def main():
    # model_dir = input("请输入模型文件夹路径: ")
    inspect_model_weights(model_dir)

if __name__ == "__main__":
    main()
相关标签
About Me
XD
Goals determine what you are going to be.
Category
标签云
多进程 Vmess FastAPI Markdown 搞笑 OpenCV Quantize Food TensorRT XML printf News Baidu Crawler Tensor Excel FP32 VGG-16 顶会 LLM Template MD5 VPN Algorithm Proxy Search InvalidArgumentError Safetensors COCO 云服务器 Ubuntu PDB QWEN Quantization CSV Qwen2.5 GPTQ 继承 HaggingFace Firewall Pillow SVR Linux Base64 NameSilo BF16 Conda ONNX diffusers torchinfo Pandas DeepStream Rebuttal PyCharm SQLite Sklearn Qwen PIP Domain NLTK Knowledge transformers 多线程 Docker BeautifulSoup AI Streamlit LoRA Mixtral Logo 公式 Diagram Zip PyTorch FP16 GIT RGB Michelin CC Ptyhon Pickle Translation Tracking 图标 Augmentation TSV GoogLeNet Heatmap Hilton FP8 YOLO Plotly Transformers EXCEL PDF VSCode Llama LLAMA WAN Bin Disk Paper uWSGI Windows Clash Review 签证 uwsgi Use v0.dev git Freesound Land Math GGML Jetson 论文速读 Bert FlashAttention Data Miniforge CAM UI Statistics scipy Website ms-swift GPT4 递归学习法 版权 Interview CUDA 第一性原理 TTS Python 图形思考法 Breakpoint Color ResNet-50 腾讯云 阿里云 Github FP64 Gemma Anaconda Django 算法题 BTC Bipartite tqdm Card 报税 Random SQL Google 论文 TensorFlow ChatGPT 净利润 SAM 证件照 JSON ModelScope CV RL llama.cpp LaTeX hf Claude DeepSeek Git Numpy 域名 XGBoost tar Password WebCrawler 飞书 Paddle 音频 Dataset logger Pytorch C++ 财报 OCR Vim OpenAI Nginx Attention Cloudreve Qwen2 Jupyter CLAP Shortcut Tiktoken SPIE Video Distillation HuggingFace mmap v2ray Plate 关于博主 Datetime NLP API CEIR Agent Image2Text RAR UNIX Input Hotel Hungarian LeetCode CTC icon Permission 强化学习 git-lfs Bitcoin Animate Web Magnet IndexTTS2
站点统计

本站现有博文332篇,共被浏览900206

本站已经建立2601天!

热门文章
文章归档
回到顶部