[object Object]

← back to Exo

fix layer calculation for sharded llama

6ee0547eff652da2c8f6b5e67deaf004341efc67 · 2024-07-13 15:39:31 -0700 · Alex Cheema

Files touched

Diff

commit 6ee0547eff652da2c8f6b5e67deaf004341efc67
Author: Alex Cheema <alexcheema123@gmail.com>
Date:   Sat Jul 13 15:39:31 2024 -0700

    fix layer calculation for sharded llama
---
 inference/mlx/models/sharded_llama.py | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/inference/mlx/models/sharded_llama.py b/inference/mlx/models/sharded_llama.py
index db67607d..c2c851b7 100644
--- a/inference/mlx/models/sharded_llama.py
+++ b/inference/mlx/models/sharded_llama.py
@@ -166,7 +166,7 @@ class LlamaModel(nn.Module):
         assert self.vocab_size > 0
         self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
         self.layers = [
-            TransformerBlock(args=args) for _ in range(args.shard.n_layers)
+            TransformerBlock(args=args) for _ in range(args.shard.end_layer - args.shard.start_layer + 1)
         ]
         self.norm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
 

← 445eda15 dynamically assign shards to nodes deterministically weighte  ·  back to Exo  ·  make StatefulShardedModel callable, add some tests for mlx s 850b72d3 →