Andhs commited on
Commit
db7c0b7
·
verified ·
1 Parent(s): 6c8447d

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +23 -7
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
- """Shallow clone of prefix KV-cache to avoid expensive copy.deepcopy."""
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
- return pkv
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=cond_past_clone, use_cache=False
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=uncond_past_clone, use_cache=False
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=cond_past_clone, use_cache=False
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=uncond_past_clone, use_cache=False
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(