18 lines
726 B
Python
18 lines
726 B
Python
"""List selected H3 checkpoint tensor shapes."""
|
|
|
|
import argparse
|
|
|
|
from safetensors import safe_open
|
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--checkpoint", default="/models/minimax_h3_fl2va_pruned_nvfp4.safetensors")
|
|
parser.add_argument("--prefix", action="append", default=[])
|
|
args = parser.parse_args()
|
|
prefixes = tuple(args.prefix) or ("adaln_t_table", "blocks.0.adaln_proj", "final_layer.adaln_proj")
|
|
|
|
with safe_open(args.checkpoint, framework="pt", device="cpu") as checkpoint:
|
|
for name in checkpoint.keys():
|
|
if any(name == prefix or name.startswith(prefix) for prefix in prefixes):
|
|
tensor = checkpoint.get_tensor(name)
|
|
print(name, tuple(tensor.shape), tensor.dtype)
|