fix batch handler for tp
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user