File size: 3,854 Bytes
31dc8dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 | import base64
import time
from typing import Any, Callable, Optional
import os
from ..types import MessageList, SamplerBase
import torch
import gc
from transformers import AutoModel, AutoTokenizer
class LLaDASampler(SamplerBase):
"""
Sample from LLaDA
"""
def __init__(
self,
model_name: str = "meta-llama/Llama-3.1-8B-Instruct",
generate_fn: Optional[Callable] = None,
generation_kwargs: Optional[dict] = None,
kv_cache_masked: Optional[bool] = False,
kv_cache_decoded: Optional[bool] = False,
):
if not kv_cache_masked and not kv_cache_decoded:
from generation_utils.llada_generate import generate as llada_ori_generate
self.model = AutoModel.from_pretrained(
model_name,
trust_remote_code=True,
torch_dtype=torch.bfloat16,
).eval().requires_grad_(False)
self.generate_fn = llada_ori_generate
elif kv_cache_decoded:
from models.modeling_llada_kv_cache import LLaDAModelLM
from generation_utils.kv_cache import generate as llada_kv_generate
self.model = LLaDAModelLM.from_pretrained(
model_name,
trust_remote_code=True,
torch_dtype=torch.bfloat16,
).eval().requires_grad_(False)
self.generate_fn = llada_kv_generate
elif kv_cache_masked:
from models.modeling_llada_qcache_improved import LLaDAModelLM
from generation_utils.q_cache_improved import generate as llada_kv_generate_masked
self.model = LLaDAModelLM.from_pretrained(
model_name,
trust_remote_code=True,
torch_dtype=torch.bfloat16,
).eval().requires_grad_(False)
self.generate_fn = llada_kv_generate_masked
else:
raise NotImplementedError()
self.tokenizer = AutoTokenizer.from_pretrained(
model_name, trust_remote_code=True
)
self.generation_kwargs = generation_kwargs
def init_model(self):
self.model = self.model.cuda()
def _handle_text(self, text: str) -> dict[str, Any]:
return {"type": "input_text", "text": text}
def _pack_message(self, role: str, content: Any) -> dict[str, Any]:
return {"role": role, "content": content}
def _free_memory(self):
del self.model
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
import time
time.sleep(10)
def __call__(self, seq_idx: int, message_list: MessageList) -> str:
with torch.inference_mode():
prompt = self.tokenizer.apply_chat_template(message_list, add_generation_prompt=True, tokenize=False)
#print(prompt)
#print(message_list)
input_ids = self.tokenizer(
[prompt],
padding_side = 'left',
padding = 'longest'
)['input_ids']
input_ids = torch.tensor(input_ids).to(self.model.device)
#set_random_seed(42)
out = self.generate_fn(
self.model, self.tokenizer, input_ids,
**self.generation_kwargs
#steps=128, gen_length=128, block_length=64,
#temperature=0., cfg_scale=0.,
#remasking='random',
#enable_cache=True,
#cache_reloading_step=4,
#window_size=args.window_size
)
res = self.tokenizer.batch_decode(
out[:, input_ids.shape[1]:],
skip_special_tokens=True
)[0]
#print(prompt)
#print(res)
return seq_idx, res
|