fix batch handler for tp

This commit is contained in:
Ryuichi Leo Takashige
2026-01-27 12:53:40 +00:00
parent 1d1256c769
commit 5e3cd73a9e
2 changed files with 26 additions and 4 deletions
+16 -1
View File
@@ -3,13 +3,28 @@
Fixes:
- UnboundLocalError on `outputs` in TemplateAPI.amodel_call when API returns error
- Prevents eval crash on transient API failures (returns None instead of raising)
- Compatibility with transformers 5.x (missing AutoModelForVision2Seq)
Usage: python -m bench.lm_eval_patched [lm_eval args...]
"""
# ruff: noqa: I001, E402
# pyright: reportMissingTypeStubs=false, reportUnknownVariableType=false
# pyright: reportUnknownMemberType=false, reportAny=false
# ruff: noqa: I001
# MUST patch transformers BEFORE any lm_eval imports
# AutoModelForVision2Seq/AutoModelForImageTextToText were removed in transformers 5.0
# Patch the lazy module's __getattr__ to return stubs for missing classes
from transformers.utils import import_utils
_original_getattr = import_utils._LazyModule.__getattr__
def _patched_getattr(self: object, name: str) -> object:
if name in ("AutoModelForVision2Seq", "AutoModelForImageTextToText"):
return type(name, (), {}) # Return a stub class
return _original_getattr(self, name) # type: ignore
import_utils._LazyModule.__getattr__ = _patched_getattr # type: ignore
import functools
from typing import Any
+10 -3
View File
@@ -62,7 +62,7 @@ from exo.shared.types.worker.runners import (
RunnerStatus,
RunnerWarmingUp,
)
from exo.shared.types.worker.shards import ShardMetadata
from exo.shared.types.worker.shards import ShardMetadata, TensorShardMetadata
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.image import (
DistributedImageModel,
@@ -298,16 +298,23 @@ def main(
# Initialize batch handler for text generation models
if BATCH_ENABLED:
# For tensor parallelism, distributed ops are handled inside model layers
# so batch handler should use world_size=1 (no pipelining)
batch_world_size = (
1
if isinstance(shard_metadata, TensorShardMetadata)
else shard_metadata.world_size
)
batch_handler = BatchedInferenceHandler(
model=model,
tokenizer=tokenizer,
model_id=shard_metadata.model_card.model_id,
device_rank=device_rank,
world_size=shard_metadata.world_size,
world_size=batch_world_size,
max_batch_size=BATCH_MAX_SIZE,
)
logger.info(
f"Batch handler initialized (max_batch_size={BATCH_MAX_SIZE}, world_size={shard_metadata.world_size})"
f"Batch handler initialized (max_batch_size={BATCH_MAX_SIZE}, world_size={batch_world_size})"
)
kv_prefix_cache = KVPrefixCache(tokenizer)