Files
2026-07-13 13:33:03 +08:00

65 lines
1.8 KiB
Python

import os
import argparse
from tqdm import tqdm
import MNN.llm as mnnllm
from datasets import load_dataset
import torch
import copy
def main(args):
model = mnnllm.create(args.mnn_path)
model.set_config({'all_logits': True, 'use_template': False})
model.set_config({'enable_debug': True})
model.load()
model.enable_collection_mode(1, args.output_path, args.target_sparsity)
eval_dataset = args.eval_dataset
dataset_parts = eval_dataset.split("/")
if len(dataset_parts) < 2:
raise ValueError("eval_dataset must be formatted as dataset/config or namespace/dataset/config.")
dataset_name = "/".join(dataset_parts[:-1])
dataset_dir = dataset_parts[-1]
dataset = load_dataset(dataset_name, dataset_dir, split="test")
input_ids = model.tokenizer_encode("\n\n".join(dataset["text"]))
input_ids = input_ids[:args.length]
_ = model.forward(input_ids)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Get thresholds from MNN model.")
parser.add_argument(
"-m",
"--mnn-path",
type=str,
required=True,
help="mnn model path",
)
parser.add_argument(
"-d", "--eval_dataset", type=str, default='Salesforce/wikitext/wikitext-2-raw-v1', help="dataset, default is `Salesforce/wikitext/wikitext-2-raw-v1`."
)
parser.add_argument(
"-o", "--output-path", type=str, default='thresholds.json', help="output path, default is `thresholds.json`."
)
parser.add_argument(
"-t", "--target-sparsity", type=float, default=0.5, help="target sparsity, default is 0.5."
)
parser.add_argument(
"-l", "--length", type=int, default=512, help="length of samples, default is 512."
)
args = parser.parse_args()
main(args)