Fix tests broken in recent commits (#1239)

We'll have good CI soon...

## Test Plan

### Automated Testing
Wroks
This commit is contained in:
rltakashige
2026-01-21 18:32:49 +00:00
committed by GitHub
parent 307f454b96
commit a354aaa3e5
2 changed files with 10 additions and 8 deletions
@@ -19,7 +19,7 @@ from exo.shared.types.tasks import ChatCompletionTaskParams
from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata
from exo.worker.engines.mlx import Model
from exo.worker.engines.mlx.generator.generate import mlx_generate
from exo.worker.engines.mlx.utils_mlx import shard_and_load
from exo.worker.engines.mlx.utils_mlx import apply_chat_template, shard_and_load
class MockLayer(nn.Module):
@@ -119,11 +119,12 @@ def run_gpt_oss_pipeline_device(
max_tokens=max_tokens,
)
prompt = apply_chat_template(tokenizer, task)
generated_text = ""
for response in mlx_generate(
model=model,
tokenizer=tokenizer,
task=task,
model=model, tokenizer=tokenizer, task=task, prompt=prompt
):
generated_text += response.text
if response.finish_reason is not None:
@@ -186,11 +187,14 @@ def run_gpt_oss_tensor_parallel_device(
max_tokens=max_tokens,
)
prompt = apply_chat_template(tokenizer, task)
generated_text = ""
for response in mlx_generate(
model=model,
tokenizer=tokenizer,
task=task,
prompt=prompt,
):
generated_text += response.text
if response.finish_reason is not None:
@@ -105,7 +105,7 @@ def event_loop():
TEST_MODELS,
)
@pytest.mark.asyncio
async def test_tokenizer_encode_decode(short_id: str, model_card: ModelCard) -> None:
async def test_tokenizer_encode_decode(model_card: ModelCard) -> None:
"""Test that tokenizer can encode and decode text correctly."""
model_id = model_card.model_id
@@ -170,9 +170,7 @@ async def test_tokenizer_encode_decode(short_id: str, model_card: ModelCard) ->
TEST_MODELS,
)
@pytest.mark.asyncio
async def test_tokenizer_has_required_attributes(
short_id: str, model_card: ModelCard
) -> None:
async def test_tokenizer_has_required_attributes(model_card: ModelCard) -> None:
"""Test that tokenizer has required attributes for inference."""
model_id = model_card.model_id