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

本站现有博文333篇,共被浏览908510

本站已经建立2610天!

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