CoolFace
Apppublic

teragron/TinyStories

sourceHugging Faceupdated 3y agoView on Hugging Face
3likes
sample.py80 linesDownload Raw Back to root
1"""2Sample from the trained model with PyTorch3"""4import os5import pickle6from contextlib import nullcontext7import torch8from model import ModelArgs, Transformer9from tokenizer import Tokenizer10 11from tinystories import get_tokenizer_model_path12 13# -----------------------------------------------------------------------------14checkpoint = 'out/ckpt.pt'15start = "" # or "<|endoftext|>" or etc. Can also specify a file, use as: "FILE:prompt.txt"16num_samples = 1 # number of samples to draw17max_new_tokens = 100 # number of tokens generated in each sample18temperature = 1.0 # 1.0 = no change, < 1.0 = less random, > 1.0 = more random, in predictions19top_k = 300 # retain only the top_k most likely tokens, clamp others to have 0 probability20tokenizer = "" # override the tokenizer model path21seed = 133722device = 'cuda' if torch.cuda.is_available() else 'cpu' # examples: 'cpu', 'cuda', 'cuda:0', 'cuda:1', etc.23#dtype = 'bfloat16' if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else 'float16' # 'float32' or 'bfloat16' or 'float16'24dtype = "float32"25compile = False # use PyTorch 2.0 to compile the model to be faster26exec(open('configurator.py').read()) # overrides from command line or config file27# -----------------------------------------------------------------------------28 29torch.manual_seed(seed)30torch.cuda.manual_seed(seed)31torch.backends.cuda.matmul.allow_tf32 = True # allow tf32 on matmul32torch.backends.cudnn.allow_tf32 = True # allow tf32 on cudnn33device_type = 'cuda' if 'cuda' in device else 'cpu' # for later use in torch.autocast34ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[dtype]35ctx = nullcontext() if device_type == 'cpu' else torch.amp.autocast(device_type=device_type, dtype=ptdtype)36 37# init from a model saved in a specific directory38checkpoint_dict = torch.load(checkpoint, map_location=device)39gptconf = ModelArgs(**checkpoint_dict['model_args'])40model = Transformer(gptconf)41state_dict = checkpoint_dict['model']42unwanted_prefix = '_orig_mod.'43for k,v in list(state_dict.items()):44    if k.startswith(unwanted_prefix):45        state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)46model.load_state_dict(state_dict, strict=False)47 48model.eval()49model.to(device)50if compile:51    print("Compiling the model...")52    model = torch.compile(model) # requires PyTorch 2.0 (optional)53 54# load the tokenizer55vocab_source = checkpoint_dict["config"].get("vocab_source", "llama2")56vocab_size = gptconf.vocab_size57if tokenizer:58    # a specific tokenizer is provided, use it59    tokenizer_model = tokenizer60else:61    # let's try to find the tokenizer model automatically. bit gross here...62    query_vocab_size = 0 if vocab_source == "llama2" else vocab_size63    tokenizer_model = get_tokenizer_model_path(vocab_size=query_vocab_size)64enc = Tokenizer(tokenizer_model=tokenizer_model)65 66# encode the beginning of the prompt67if start.startswith('FILE:'):68    with open(start[5:], 'r', encoding='utf-8') as f:69        start = f.read()70start_ids = enc.encode(start, bos=True, eos=False)71x = (torch.tensor(start_ids, dtype=torch.long, device=device)[None, ...])72 73# run generation74with torch.no_grad():75    with ctx:76        for k in range(num_samples):77            y = model.generate(x, max_new_tokens, temperature=temperature, top_k=top_k)78            print(enc.decode(y[0].tolist()))79            print('---------------')80