pipenetwork commited on
Commit
2665d43
·
verified ·
1 Parent(s): d8dc477

Add step_callback hook

Browse files
Files changed (1) hide show
  1. nemotron_twotower_mlx.py +5 -0
nemotron_twotower_mlx.py CHANGED
@@ -276,6 +276,7 @@ class TwoTowerModel(nn.Module):
276
  def generate_mask_diffusion(
277
  self, input_ids, max_new_tokens=128, block_size=16, steps_per_block=16,
278
  mask_token_id=3, confidence_threshold=0.9, eos_token_id=None, verbose=False,
 
279
  ):
280
  assert max_new_tokens % block_size == 0
281
  B = input_ids.shape[0]
@@ -286,6 +287,8 @@ class TwoTowerModel(nn.Module):
286
 
287
  for blk in range(num_blocks):
288
  xt = mx.full((B, block_size), mask_token_id, dtype=mx.int32)
 
 
289
  for step in range(steps_per_block):
290
  is_masked = (xt == mask_token_id)
291
  n_masked = int(is_masked.sum().item())
@@ -328,6 +331,8 @@ class TwoTowerModel(nn.Module):
328
  new_xt.append(row)
329
  xt = mx.stack(new_xt)
330
  mx.eval(xt)
 
 
331
 
332
  context_ids = mx.concatenate([context_ids, xt], axis=1)
333
  caches = self.extend_context_cache(xt, caches)
 
276
  def generate_mask_diffusion(
277
  self, input_ids, max_new_tokens=128, block_size=16, steps_per_block=16,
278
  mask_token_id=3, confidence_threshold=0.9, eos_token_id=None, verbose=False,
279
+ step_callback=None,
280
  ):
281
  assert max_new_tokens % block_size == 0
282
  B = input_ids.shape[0]
 
287
 
288
  for blk in range(num_blocks):
289
  xt = mx.full((B, block_size), mask_token_id, dtype=mx.int32)
290
+ if step_callback is not None:
291
+ step_callback(blk, -1, xt, context_ids)
292
  for step in range(steps_per_block):
293
  is_masked = (xt == mask_token_id)
294
  n_masked = int(is_masked.sum().item())
 
331
  new_xt.append(row)
332
  xt = mx.stack(new_xt)
333
  mx.eval(xt)
334
+ if step_callback is not None:
335
+ step_callback(blk, step, xt, context_ids)
336
 
337
  context_ids = mx.concatenate([context_ids, xt], axis=1)
338
  caches = self.extend_context_cache(xt, caches)