← back to Exo
also initialize embed_tokens if last layer and tie_word_embeddings true
ad09b4b3d9d8999bbedfd777769f94118b833e6b · 2024-10-10 14:11:11 -0700 · Alex Cheema
Files touched
M exo/inference/mlx/models/llama.py
Diff
commit ad09b4b3d9d8999bbedfd777769f94118b833e6b
Author: Alex Cheema <alexcheema123@gmail.com>
Date: Thu Oct 10 14:11:11 2024 -0700
also initialize embed_tokens if last layer and tie_word_embeddings true
---
exo/inference/mlx/models/llama.py | 8 ++++----
1 file changed, 4 insertions(+), 4 deletions(-)
diff --git a/exo/inference/mlx/models/llama.py b/exo/inference/mlx/models/llama.py
index 11b6716a..8b069ecc 100644
--- a/exo/inference/mlx/models/llama.py
+++ b/exo/inference/mlx/models/llama.py
@@ -32,15 +32,15 @@ class LlamaModel(nn.Module):
self.vocab_size = args.vocab_size
self.num_hidden_layers = args.num_hidden_layers
assert self.vocab_size > 0
- if self.args.shard.is_first_layer():
+ if args.shard.is_first_layer() or (args.shard.is_last_layer() and args.tie_word_embeddings):
self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size)
self.layers = []
for i in range(self.num_hidden_layers):
- if self.args.shard.start_layer <= i <= self.args.shard.end_layer:
+ if args.shard.start_layer <= i <= args.shard.end_layer:
self.layers.append(TransformerBlock(args=args))
else:
self.layers.append(IdentityBlock())
- if self.args.shard.is_last_layer():
+ if args.shard.is_last_layer():
self.norm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
def __call__(
@@ -74,7 +74,7 @@ class Model(nn.Module):
self.args = args
self.model_type = args.model_type
self.model = LlamaModel(args)
- if self.args.shard.is_last_layer():
+ if args.shard.is_last_layer():
if not args.tie_word_embeddings:
self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False)
← fbc407c6 make llama-3.2-1b the default for tests so they run faster
·
back to Exo
·
run unit test on llama 3.2 1b for faster test ae74d2da →