"""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)