Spaces:
Sleeping
Sleeping
Upload app.py
Browse files
app.py
CHANGED
|
@@ -75,16 +75,30 @@ def build_staircase_attention_mask(x, block_size, pad_id):
|
|
| 75 |
|
| 76 |
|
| 77 |
def clone_past_key_values(pkv):
|
| 78 |
-
"""
|
| 79 |
if pkv is None:
|
| 80 |
return None
|
|
|
|
| 81 |
if isinstance(pkv, tuple):
|
| 82 |
return tuple(
|
| 83 |
(k.clone() if k is not None else None, v.clone() if v is not None else None)
|
| 84 |
for k, v in pkv
|
| 85 |
)
|
| 86 |
-
|
| 87 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 88 |
|
| 89 |
def diffusion_step_block(logits, x_block, mask_block, num_transfer, temperature, remasking):
|
| 90 |
"""Vectorized diffusion step — no per-sample Python loops."""
|
|
@@ -172,18 +186,20 @@ def generate(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size
|
|
| 172 |
# Clone prefix caches once per block instead of every step
|
| 173 |
cond_past_clone = clone_past_key_values(cond_past)
|
| 174 |
uncond_past_clone = clone_past_key_values(uncond_past) if uncond_past is not None else None
|
|
|
|
| 175 |
for t in range(eff_steps):
|
| 176 |
x_blk = x[:, T_prefix:T_total]
|
| 177 |
m_blk = x_blk == mask_id
|
| 178 |
cond_logits = model(
|
| 179 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 180 |
-
past_key_values=
|
| 181 |
).logits
|
| 182 |
logits = cond_logits
|
|
|
|
| 183 |
if cfg_scale > 0:
|
| 184 |
un_logits = model(
|
| 185 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 186 |
-
past_key_values=
|
| 187 |
).logits
|
| 188 |
logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
|
| 189 |
x_blk_new = diffusion_step_block(
|
|
@@ -267,13 +283,13 @@ def generate_stream(model, tokenizer, prompt, steps=128, max_new_tokens=128, blo
|
|
| 267 |
m_blk = x_blk == mask_id
|
| 268 |
cond_logits = model(
|
| 269 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 270 |
-
past_key_values=
|
| 271 |
).logits
|
| 272 |
logits = cond_logits
|
| 273 |
if cfg_scale > 0:
|
| 274 |
un_logits = model(
|
| 275 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 276 |
-
past_key_values=
|
| 277 |
).logits
|
| 278 |
logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
|
| 279 |
x_blk_new = diffusion_step_block(
|
|
|
|
| 75 |
|
| 76 |
|
| 77 |
def clone_past_key_values(pkv):
|
| 78 |
+
"""Clone KV-cache. Fast path for tuples and Cache objects; falls back to deepcopy."""
|
| 79 |
if pkv is None:
|
| 80 |
return None
|
| 81 |
+
# Fast path: legacy tuple format
|
| 82 |
if isinstance(pkv, tuple):
|
| 83 |
return tuple(
|
| 84 |
(k.clone() if k is not None else None, v.clone() if v is not None else None)
|
| 85 |
for k, v in pkv
|
| 86 |
)
|
| 87 |
+
# Fast path: transformers Cache objects (DynamicCache, etc.)
|
| 88 |
+
if hasattr(pkv, 'key_cache') and hasattr(pkv, 'value_cache'):
|
| 89 |
+
try:
|
| 90 |
+
new_cache = pkv.__class__()
|
| 91 |
+
new_cache.key_cache = [k.clone() for k in pkv.key_cache]
|
| 92 |
+
new_cache.value_cache = [v.clone() for v in pkv.value_cache]
|
| 93 |
+
for attr in ('_seen_tokens', 'seen_tokens'):
|
| 94 |
+
if hasattr(pkv, attr):
|
| 95 |
+
setattr(new_cache, attr, getattr(pkv, attr))
|
| 96 |
+
return new_cache
|
| 97 |
+
except Exception:
|
| 98 |
+
pass
|
| 99 |
+
# Fallback
|
| 100 |
+
import copy
|
| 101 |
+
return copy.deepcopy(pkv)
|
| 102 |
|
| 103 |
def diffusion_step_block(logits, x_block, mask_block, num_transfer, temperature, remasking):
|
| 104 |
"""Vectorized diffusion step — no per-sample Python loops."""
|
|
|
|
| 186 |
# Clone prefix caches once per block instead of every step
|
| 187 |
cond_past_clone = clone_past_key_values(cond_past)
|
| 188 |
uncond_past_clone = clone_past_key_values(uncond_past) if uncond_past is not None else None
|
| 189 |
+
|
| 190 |
for t in range(eff_steps):
|
| 191 |
x_blk = x[:, T_prefix:T_total]
|
| 192 |
m_blk = x_blk == mask_id
|
| 193 |
cond_logits = model(
|
| 194 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 195 |
+
past_key_values=clone_past_key_values(cond_past), use_cache=False
|
| 196 |
).logits
|
| 197 |
logits = cond_logits
|
| 198 |
+
|
| 199 |
if cfg_scale > 0:
|
| 200 |
un_logits = model(
|
| 201 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 202 |
+
past_key_values=clone_past_key_values(uncond_past), use_cache=False
|
| 203 |
).logits
|
| 204 |
logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
|
| 205 |
x_blk_new = diffusion_step_block(
|
|
|
|
| 283 |
m_blk = x_blk == mask_id
|
| 284 |
cond_logits = model(
|
| 285 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 286 |
+
past_key_values=clone_past_key_values(cond_past), use_cache=False
|
| 287 |
).logits
|
| 288 |
logits = cond_logits
|
| 289 |
if cfg_scale > 0:
|
| 290 |
un_logits = model(
|
| 291 |
x_blk, attention_mask=attn_blk, position_ids=pos_blk,
|
| 292 |
+
past_key_values=clone_past_key_values(uncond_past), use_cache=False
|
| 293 |
).logits
|
| 294 |
logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
|
| 295 |
x_blk_new = diffusion_step_block(
|