DeepSeek-Flash-Mini / load_and_generate.py
nowordsxiaomu's picture
Initial release: DeepSeek-Flash-Mini nano (15M MoE, MLA+MTP)
5e6d9f5 verified
Raw
History Blame Contribute Delete
1.8 kB
"""Self-contained loader for DeepSeek-Flash-Mini (HF export).
Usage:
from load_and_generate import load_model
model, cfg = load_model(".") # repo dir containing config.json + model.safetensors
# or CLI:
python load_and_generate.py --prompt "Once upon a time" --max-new-tokens 80
"""
import argparse, json
import torch
from safetensors.torch import load_file
from config import ModelConfig
from model import DeepSeekFlashMini
from dataio.tokenizer import load_tokenizer
from generate import Generator
def load_model(repo_dir: str = "."):
with open(f"{repo_dir}/config.json", encoding="utf-8") as f:
cfg = ModelConfig(**json.load(f))
model = DeepSeekFlashMini(cfg)
sd = load_file(f"{repo_dir}/model.safetensors")
model.load_state_dict(sd)
model.eval()
return model, cfg
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--repo-dir", default=".")
ap.add_argument("--prompt", default="Once upon a time")
ap.add_argument("--max-new-tokens", type=int, default=80)
ap.add_argument("--temperature", type=float, default=0.8)
ap.add_argument("--top-k", type=int, default=40)
ap.add_argument("--top-p", type=float, default=0.9)
ap.add_argument("--device", default="cpu")
ap.add_argument("--spec", action="store_true", help="use MTP speculative decoding")
args = ap.parse_args()
model, cfg = load_model(args.repo_dir)
tok = load_tokenizer(f"{args.repo_dir}/tokenizer.json")
gen = Generator(model, tok, device=args.device)
text = gen.generate(args.prompt, max_new_tokens=args.max_new_tokens,
temperature=args.temperature, top_k=args.top_k,
top_p=args.top_p, speculative=args.spec)
print(text)
if __name__ == "__main__":
main()