Instructions to use pipenetwork/Nemotron-Labs-TwoTower-30B-A3B-mlx-6bit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use pipenetwork/Nemotron-Labs-TwoTower-30B-A3B-mlx-6bit with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir Nemotron-Labs-TwoTower-30B-A3B-mlx-6bit pipenetwork/Nemotron-Labs-TwoTower-30B-A3B-mlx-6bit
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Add step_callback hook
Browse files- 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)
|