File size: 11,800 Bytes
3a464db | 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 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 | import torch
import logging
from dinfer import DiffusionLLM
from dinfer.decoding.utils import TokenArray
logger = logging.getLogger(__name__)
class IterSmoothDiffusionLLM(DiffusionLLM):
""" This diffusion LLM inference generates tokens block by block.
The decoding algorithm break the generation sequence into blocks.
It runs diffusion iterations on the first block and decodes all tokens
in the block before moving to the next block.
This is a classifical dLLM decoding algorithm.
"""
def __init__(self, model, decoder, iterator_factory, early_stop=True, cache_factory=None, maximum_unroll=4, expected_tpf=8,
cont_weight=0.3, cont_weight_init=0.15, cont_weight_growth=0.02, threshold_decay=0.02):
self.model = model
self.cache_factory = cache_factory
self.decoder = decoder
self.iterator_factory = iterator_factory
self.num_forwards = 0
self.cache_updates = 0
self.early_stop = early_stop
self.cont_weight = cont_weight
if isinstance(self.model, torch.nn.parallel.DistributedDataParallel):
self.h2e = self.model.module.h2e
else:
self.h2e = self.model.h2e
self.cont_weight_init = cont_weight_init
self.cont_weight_growth = cont_weight_growth
self.threshold_decay = threshold_decay
self.maximum_unroll = maximum_unroll
self.expected_tpf = expected_tpf
@ torch.no_grad()
def generate(self, prompt, gen_length=128, block_length=128):
''' Generate tokens with diffusion iterations block by block.
'''
x = TokenArray(prompt, gen_length, self.decoder.mask_id, self.decoder.eos_id, self.model.device)
it = self.iterator_factory.create(x, block_length)
inputs_embeds = self.h2e(x.data)
iter_no = 0
kv_cache = self.cache_factory.create() if self.cache_factory is not None else None
for block_id, (block_loc, block) in enumerate(it):
self.decoder.block_init(block, block_id)
while (block == self.decoder.mask_id).sum() > 0:
unroll_k = max(min((block == self.decoder.mask_id).sum()//self.expected_tpf, self.maximum_unroll), 1)
for unroll_i in range(unroll_k):
iter_cont_weight = min(self.cont_weight_init+self.cont_weight_growth*iter_no, self.cont_weight)
iter_threshold = max(1-iter_no*self.threshold_decay, self.decoder.threshold)
# Update KV-cache
if kv_cache is not None and kv_cache.require_update(iter_no, block_loc.start, block_loc.end):
output = self.model(inputs_embeds=inputs_embeds, use_cache=True)
self.num_forwards += 1
# use the generated output to decode.
self.decoder.decode(output.logits[:, block_loc.start:block_loc.end], block_loc.start, block_loc.end, x, iter_threshold)
# update KV-cache
mask_index = (x.data == self.decoder.mask_id)
inputs_embeds = self.h2e(x.data, mask_index, output.logits, iter_cont_weight)
kv_cache.update(output.past_key_values)
past_key_values, replace_position = kv_cache.get_key_values(block_loc.start, block_loc.end)
self.cache_updates += 1
iter_no += 1
iter_cont_weight = min(self.cont_weight_init+self.cont_weight_growth*iter_no, self.cont_weight)
iter_threshold = max(1-iter_no*self.threshold_decay, self.decoder.threshold)
if kv_cache is None:
logits = self.model(inputs_embeds=inputs_embeds).logits
self.decoder.decode(logits[:, block_loc.start:block_loc.end], block_loc.start, block_loc.end, x, iter_threshold)
mask_index = (x.data == self.decoder.mask_id)
inputs_embeds = self.h2e(x.data, mask_index, logits, iter_cont_weight)
elif kv_cache.cache_type == 'prefix':
logits = self.model(inputs_embeds=inputs_embeds[:, block_loc.start:], past_key_values=past_key_values, use_cache=True,
replace_position=replace_position).logits
block_length = block_loc.end - block_loc.start
self.decoder.decode(logits[:, :block_length], block_loc.start, block_loc.end, x, iter_threshold)
mask_index = (x.data[:, block_loc.start:] == self.decoder.mask_id)
inputs_embeds[:, block_loc.start:] = self.h2e(x.data[:, block_loc.start:], mask_index, logits, iter_cont_weight)
else:
# cache position is the position between current_block_start and current_block_end
logits = self.model(inputs_embeds=inputs_embeds[:, block_loc.start:block_loc.end], past_key_values=past_key_values, use_cache=True,
replace_position=replace_position).logits
self.decoder.decode(logits, block_loc.start, block_loc.end, x, iter_threshold)
mask_index = (x.data[:, block_loc.start:block_loc.end] == self.decoder.mask_id)
inputs_embeds[:, block_loc.start:block_loc.end] = self.h2e(x.data[:, block_loc.start:block_loc.end], mask_index, logits, iter_cont_weight)
self.num_forwards += 1
iter_no += 1
if self.early_stop and torch.any(x[:, block_loc.start:block_loc.end] == self.decoder.eos_id):
# Find the first location of EOS and set all tokens after the location to EOS.
# Here we assume that don't perform remasking.
# TODO(zhengda) here we assume the batch size is 1.
x[:, block_loc.end:] = self.decoder.eos_id
break
logger.info(f'The number of diffusion iterations: {self.num_forwards}')
return x.get_generated_tokens()
class IterSmoothWithVicinityCacheDiffusionLLM(DiffusionLLM):
""" This diffusion LLM inference generates tokens with vicinity cache and iteration smoothing.
"""
def __init__(self, model, decoder, iterator_factory, cache_factory, maximum_unroll=4, expected_tpf=8,
prefix_look=0, after_look=0, warmup_steps=0, early_stop=True, cont_weight=0.3,
cont_weight_init=0.15, cont_weight_growth=0.02, threshold_decay=0.02):
self.model = model
self.cache_factory = cache_factory
self.decoder = decoder
self.iterator_factory = iterator_factory
self.num_forwards = 0
self.cache_updates = 0
self.prefix_look = int(prefix_look)
self.after_look = int(after_look)
self.warmup_steps = int(warmup_steps)
self.early_stop = early_stop
self.cont_weight = cont_weight
if isinstance(self.model, torch.nn.parallel.DistributedDataParallel):
self.h2e = self.model.module.h2e
else:
self.h2e = self.model.h2e
self.cont_weight_init = cont_weight_init
self.cont_weight_growth = cont_weight_growth
self.threshold_decay = threshold_decay
self.maximum_unroll = maximum_unroll
self.expected_tpf = expected_tpf
assert cache_factory is not None, "This class requires a KV-cache."
@ torch.no_grad()
def generate(self, prompt, gen_length=128, block_length=128):
''' Generate tokens with diffusion iterations block by block.
'''
x = TokenArray(prompt, gen_length, self.decoder.mask_id, self.decoder.eos_id, self.model.device)
it = self.iterator_factory.create(x, block_length)
kv_cache = self.cache_factory.create()
prompt_len = x.prompt.shape[1]
total_len = x.total_length
inputs_embeds = self.h2e(x.data)
iter_no = 0
for block_idx, (block_loc, block) in enumerate(it):
block_start, block_end = block_loc.start, block_loc.end
left_start = max(0, block_start - self.prefix_look)
right_end = min(total_len, block_end + self.after_look)
iter_in_block = 0
while (x[:, block_start:block_end] == self.decoder.mask_id).sum() > 0:
unroll_k = max(min((block == self.decoder.mask_id).sum()//self.expected_tpf, self.maximum_unroll), 1)
for unroll_i in range(unroll_k):
iter_cont_weight = min(self.cont_weight_init+self.cont_weight_growth*iter_no, self.cont_weight)
iter_threshold = max(1-iter_no*self.threshold_decay, self.decoder.threshold)
if block_idx == 0 and iter_in_block < self.warmup_steps:
out_full = self.model(inputs_embeds=inputs_embeds)
self.num_forwards += 1
self.decoder.decode(out_full.logits[:, block_start:block_end], block_start, block_end, x, iter_threshold)
mask_index = (x.data == self.decoder.mask_id)
inputs_embeds = self.h2e(x.data, mask_index, out_full.logits, iter_cont_weight)
self.cache_updates += 1
iter_in_block += 1
iter_no += 1
continue
if kv_cache.past_key_values is None or (kv_cache.require_update(iter_no, block_start, block_end) and block_idx > 0):
out_full = self.model(inputs_embeds=inputs_embeds, use_cache=True)
self.num_forwards += 1
self.decoder.decode(out_full.logits[:, block_start:block_end], block_start, block_end, x, iter_threshold)
mask_index = (x.data == self.decoder.mask_id)
inputs_embeds = self.h2e(x.data, mask_index, out_full.logits, iter_cont_weight)
kv_cache.update(out_full.past_key_values)
self.cache_updates += 1
iter_in_block+=1
iter_no += 1
#continue
iter_cont_weight = min(self.cont_weight_init+self.cont_weight_growth*iter_no, self.cont_weight)
iter_threshold = max(1-iter_no*self.threshold_decay, self.decoder.threshold)
past_key_values, replace_position = kv_cache.get_key_values(left_start, right_end)
out_step = self.model(
inputs_embeds=inputs_embeds[:, left_start:right_end],
past_key_values=past_key_values,
use_cache=True,
replace_position=replace_position
)
self.num_forwards += 1
iter_no += 1
offset = block_start - left_start
logits_block = out_step.logits[:, offset:offset + (block_end - block_start)]
self.decoder.decode(logits_block, block_start, block_end, x, iter_threshold)
mask_index = (x.data[:, left_start:right_end] == self.decoder.mask_id)
inputs_embeds[:, left_start:right_end] = self.h2e(x.data[:, left_start:right_end], mask_index, out_step.logits, iter_cont_weight)
iter_in_block += 1
if self.early_stop and torch.any(x[:, block_start:block_end] == self.decoder.eos_id):
x[:, block_end:] = self.decoder.eos_id
break
logger.info(f'The number of diffusion iterations with kv-cache: {self.num_forwards}')
return x.get_generated_tokens()
|