← back to Exo
fix kimi eos token ids
d793f5f96c38093bd0193e4267ae56e8c538e113 · 2025-11-13 10:39:14 -0800 · Alex Cheema
Files touched
M .github/configs/bench_simple.yamlM .github/workflows/bench.ymlM src/exo/engines/mlx/utils_mlx.py
Diff
commit d793f5f96c38093bd0193e4267ae56e8c538e113
Author: Alex Cheema <41707476+AlexCheema@users.noreply.github.com>
Date: Thu Nov 13 10:39:14 2025 -0800
fix kimi eos token ids
---
.github/configs/bench_simple.yaml | 44 +++++++++++++++++++--------------------
.github/workflows/bench.yml | 1 +
src/exo/engines/mlx/utils_mlx.py | 13 ++++++++++--
3 files changed, 34 insertions(+), 24 deletions(-)
diff --git a/.github/configs/bench_simple.yaml b/.github/configs/bench_simple.yaml
index 91c85020..18f7042b 100644
--- a/.github/configs/bench_simple.yaml
+++ b/.github/configs/bench_simple.yaml
@@ -43,41 +43,41 @@ stages:
# generation_length: 10
# time_between_requests: 2.0
# iterations: 5
- - name: "pp64_g64"
- prompt_length: 64
- generation_length: 64
- time_between_requests: 2.0
- iterations: 5
+ # - name: "pp64_g64"
+ # prompt_length: 64
+ # generation_length: 64
+ # time_between_requests: 2.0
+ # iterations: 5
# - name: "pp64_g512"
# prompt_length: 64
# generation_length: 512
# time_between_requests: 2.0
# iterations: 10
- - name: "pp256_g64"
- prompt_length: 256
- generation_length: 64
- time_between_requests: 2.0
- iterations: 5
+ # - name: "pp256_g64"
+ # prompt_length: 256
+ # generation_length: 64
+ # time_between_requests: 2.0
+ # iterations: 5
# - name: "pp256_g512"
# prompt_length: 256
# generation_length: 512
# time_between_requests: 2.0
# iterations: 10
- - name: "pp1024_g64"
- prompt_length: 1024
- generation_length: 64
- time_between_requests: 2.0
- iterations: 5
+ # - name: "pp1024_g64"
+ # prompt_length: 1024
+ # generation_length: 64
+ # time_between_requests: 2.0
+ # iterations: 5
# - name: "pp1024_g512"
# prompt_length: 1024
# generation_length: 512
# time_between_requests: 2.0
# iterations: 10
- - name: "pp2048_g64"
- prompt_length: 2048
- generation_length: 64
- time_between_requests: 2.0
- iterations: 5
+ # - name: "pp2048_g64"
+ # prompt_length: 2048
+ # generation_length: 64
+ # time_between_requests: 2.0
+ # iterations: 5
# - name: "pp2048_g512"
# prompt_length: 2048
# generation_length: 512
@@ -87,7 +87,7 @@ stages:
prompt_length: 4096
generation_length: 64
time_between_requests: 2.0
- iterations: 5
+ iterations: 4
# - name: "pp4096_g512"
# prompt_length: 4096
# generation_length: 512
@@ -97,7 +97,7 @@ stages:
prompt_length: 8192
generation_length: 64
time_between_requests: 2.0
- iterations: 5
+ iterations: 4
# - name: "pp8192_g512"
# prompt_length: 8192
# generation_length: 512
diff --git a/.github/workflows/bench.yml b/.github/workflows/bench.yml
index baa0d20d..dda16435 100644
--- a/.github/workflows/bench.yml
+++ b/.github/workflows/bench.yml
@@ -4,6 +4,7 @@ on: [push]
jobs:
plan:
+ if: contains(github.event.head_commit.message, '/bench')
runs-on: ubuntu-latest
outputs:
matrix: ${{ steps.build.outputs.matrix }}
diff --git a/src/exo/engines/mlx/utils_mlx.py b/src/exo/engines/mlx/utils_mlx.py
index c82fbee6..5f42ca9c 100644
--- a/src/exo/engines/mlx/utils_mlx.py
+++ b/src/exo/engines/mlx/utils_mlx.py
@@ -149,7 +149,10 @@ def initialize_mlx(
tokenizer = cast(
TokenizerWrapper,
load_tokenizer(
- model_path, tokenizer_config_extra={"trust_remote_code": True}
+ model_path,
+ tokenizer_config_extra={"trust_remote_code": True},
+ # TODO: HACK for Kimi K2 wrong eos token id
+ eos_token_ids=[163586] if "kimi-k2" in bound_instance.bound_shard().model_meta.model_id.lower() else None,
),
)
assert isinstance(tokenizer, TokenizerWrapper)
@@ -177,7 +180,13 @@ def shard_and_load(
# TODO: we should really make this opt-in, but Kimi requires trust_remote_code=True
tokenizer = cast(
TokenizerWrapper,
- load_tokenizer(model_path, tokenizer_config_extra={"trust_remote_code": True}),
+ # TODO: HACK for Kimi K2 wrong eos token id
+ load_tokenizer(
+ model_path,
+ tokenizer_config_extra={"trust_remote_code": True},
+ # TODO: HACK for Kimi K2 wrong eos token id
+ eos_token_ids=[163586] if "kimi-k2" in shard_metadata.model_meta.model_id.lower() else None,
+ ),
)
logger.info(f"Group size: {group.size()}, group rank: {group.rank()}")
← b62f6847 improved master error handling
·
back to Exo
·
Demo 28a91787 →