Instructions to use Duke-CEI-SVD/traj-mc with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Duke-CEI-SVD/traj-mc with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("Duke-CEI-SVD/traj-mc", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download code/eval/generate.py from Duke-CEI-SVD/traj-mc: direct link, hf CLI and curl.
- Browser
- Download file 7.72 kB
-
https://huggingface.co/Duke-CEI-SVD/traj-mc/resolve/main/code/eval/generate.py
- Command line
-
hf download hf://Duke-CEI-SVD/traj-mc/code/eval/generate.py
-
curl -L -o generate.py https://huggingface.co/Duke-CEI-SVD/traj-mc/resolve/main/code/eval/generate.py
7.72 kB
| import torch | |
| import numpy as np | |
| import torch.nn.functional as F | |
| from transformers import AutoTokenizer, AutoModel | |
| def add_gumbel_noise(logits, temperature): | |
| ''' | |
| The Gumbel max is a method for sampling categorical distributions. | |
| According to arXiv:2409.02908, for MDM, low-precision Gumbel Max improves perplexity score but reduces generation quality. | |
| Thus, we use float64. | |
| ''' | |
| if temperature == 0: | |
| return logits | |
| logits = logits.to(torch.float64) | |
| noise = torch.rand_like(logits, dtype=torch.float64) | |
| gumbel_noise = (- torch.log(noise)) ** temperature | |
| return logits.exp() / gumbel_noise | |
| def get_num_transfer_tokens(mask_index, steps): | |
| ''' | |
| In the reverse process, the interval [0, 1] is uniformly discretized into steps intervals. | |
| Furthermore, because LLaDA employs a linear noise schedule (as defined in Eq. (8)), | |
| the expected number of tokens transitioned at each step should be consistent. | |
| This function is designed to precompute the number of tokens that need to be transitioned at each step. | |
| ''' | |
| mask_num = mask_index.sum(dim=1, keepdim=True) | |
| base = mask_num // steps | |
| remainder = mask_num % steps | |
| num_transfer_tokens = torch.zeros(mask_num.size(0), steps, device=mask_index.device, dtype=torch.int64) + base | |
| for i in range(mask_num.size(0)): | |
| num_transfer_tokens[i, :remainder[i]] += 1 | |
| return num_transfer_tokens | |
| def generate(model, prompt, steps=128, gen_length=128, block_length=128, temperature=0., | |
| cfg_scale=0., remasking='low_confidence', mask_id=126336, | |
| eos_id=None, eot_id=None, logits_eos_inf=False, confidence_eos_eot_inf=False): | |
| ''' | |
| Args: | |
| model: Mask predictor. | |
| prompt: A tensor of shape (1, L). | |
| steps: Sampling steps, less than or equal to gen_length. | |
| gen_length: Generated answer length. | |
| block_length: Block length, less than or equal to gen_length. If less than gen_length, it means using semi_autoregressive remasking. | |
| temperature: Categorical distribution sampling temperature. | |
| cfg_scale: Unsupervised classifier-free guidance scale. | |
| remasking: Remasking strategy. 'low_confidence' or 'random'. | |
| mask_id: The toke id of [MASK] is 126336. | |
| eos_id / eot_id: ids of <|endoftext|> (126081) / <|eot_id|> (126348, Instruct only). | |
| Only consulted when one of the two switches below is on. | |
| logits_eos_inf: force the EOS logit to -inf, so EOS can never be predicted. | |
| confidence_eos_eot_inf: force the CONFIDENCE of every position whose prediction is | |
| EOS/EOT to -inf, so low_confidence remasking never commits it early. The token | |
| can still land there on the final forced steps -- this defers EOS, it does not | |
| forbid it (that is what logits_eos_inf is for). | |
| [trajmc_main ADDITION -- LLaDA-8B-Instruct only] | |
| The two EOS switches come from the official evaluation/EVAL.md Instruct table, which | |
| sets logits_eos_inf on HumanEval and confidence_eos_eot_inf on GSM8K/Math/GPQA/MBPP. | |
| They exist because the SFT data is heavily |EOS|-padded, so Instruct over-produces EOS | |
| and truncates its own answer. NOTE: no reference implementation ships with the LLaDA | |
| repo -- the paper's Instruct numbers come from the authors' internal toolkit and | |
| OpenCompass, not lm-eval -- so the semantics above are our reading of the flag names | |
| plus the stated purpose, and must be validated against the official Dense numbers | |
| before any compressed arm is trusted. | |
| Both default to False, so every Base call path is numerically unchanged. | |
| ''' | |
| if logits_eos_inf and eos_id is None: | |
| raise ValueError("logits_eos_inf=True requires eos_id") | |
| if confidence_eos_eot_inf and eos_id is None and eot_id is None: | |
| raise ValueError("confidence_eos_eot_inf=True requires eos_id and/or eot_id") | |
| x = torch.full((1, prompt.shape[1] + gen_length), mask_id, dtype=torch.long).to(model.device) | |
| x[:, :prompt.shape[1]] = prompt.clone() | |
| prompt_index = (x != mask_id) | |
| assert gen_length % block_length == 0 | |
| num_blocks = gen_length // block_length | |
| assert steps % num_blocks == 0 | |
| steps = steps // num_blocks | |
| for num_block in range(num_blocks): | |
| block_mask_index = (x[:, prompt.shape[1] + num_block * block_length: prompt.shape[1] + (num_block + 1) * block_length:] == mask_id) | |
| num_transfer_tokens = get_num_transfer_tokens(block_mask_index, steps) | |
| for i in range(steps): | |
| mask_index = (x == mask_id) | |
| if cfg_scale > 0.: | |
| un_x = x.clone() | |
| un_x[prompt_index] = mask_id | |
| x_ = torch.cat([x, un_x], dim=0) | |
| logits = model(x_).logits | |
| logits, un_logits = torch.chunk(logits, 2, dim=0) | |
| logits = un_logits + (cfg_scale + 1) * (logits - un_logits) | |
| else: | |
| logits = model(x).logits | |
| if logits_eos_inf: | |
| # Applied BEFORE both argmax and softmax so the prediction and its | |
| # confidence agree; -inf survives the temperature=0 gumbel no-op. | |
| logits[..., eos_id] = -np.inf | |
| logits_with_noise = add_gumbel_noise(logits, temperature=temperature) | |
| x0 = torch.argmax(logits_with_noise, dim=-1) # b, l | |
| if remasking == 'low_confidence': | |
| p = F.softmax(logits, dim=-1) | |
| x0_p = torch.squeeze( | |
| torch.gather(p, dim=-1, index=torch.unsqueeze(x0, -1)), -1) # b, l | |
| elif remasking == 'random': | |
| x0_p = torch.rand((x0.shape[0], x0.shape[1]), device=x0.device) | |
| else: | |
| raise NotImplementedError(remasking) | |
| if confidence_eos_eot_inf: | |
| is_eos = torch.zeros_like(x0, dtype=torch.bool) | |
| if eos_id is not None: | |
| is_eos |= (x0 == eos_id) | |
| if eot_id is not None: | |
| is_eos |= (x0 == eot_id) | |
| x0_p = x0_p.masked_fill(is_eos, -np.inf) | |
| x0_p[:, prompt.shape[1] + (num_block + 1) * block_length:] = -np.inf | |
| x0 = torch.where(mask_index, x0, x) | |
| confidence = torch.where(mask_index, x0_p, -np.inf) | |
| transfer_index = torch.zeros_like(x0, dtype=torch.bool, device=x0.device) | |
| for j in range(confidence.shape[0]): | |
| _, select_index = torch.topk(confidence[j], k=num_transfer_tokens[j, i]) | |
| transfer_index[j, select_index] = True | |
| x[transfer_index] = x0[transfer_index] | |
| return x | |
| def main(): | |
| device = 'cuda' | |
| model = AutoModel.from_pretrained('GSAI-ML/LLaDA-8B-Instruct', trust_remote_code=True, torch_dtype=torch.bfloat16).to(device).eval() | |
| tokenizer = AutoTokenizer.from_pretrained('GSAI-ML/LLaDA-8B-Instruct', trust_remote_code=True) | |
| prompt = "Lily can run 12 kilometers per hour for 4 hours. After that, she runs 6 kilometers per hour. How many kilometers can she run in 8 hours?" | |
| # Add special tokens for the Instruct model. The Base model does not require the following two lines. | |
| m = [{"role": "user", "content": prompt}, ] | |
| prompt = tokenizer.apply_chat_template(m, add_generation_prompt=True, tokenize=False) | |
| input_ids = tokenizer(prompt)['input_ids'] | |
| input_ids = torch.tensor(input_ids).to(device).unsqueeze(0) | |
| out = generate(model, input_ids, steps=128, gen_length=128, block_length=32, temperature=0., cfg_scale=0., remasking='low_confidence') | |
| print(tokenizer.batch_decode(out[:, input_ids.shape[1]:], skip_special_tokens=True)[0]) | |
| if __name__ == '__main__': | |
| main() | |