diff --git a/deepspeed/inference/config.py b/deepspeed/inference/config.py index 2ef43cb0239d..25994a110d31 100644 --- a/deepspeed/inference/config.py +++ b/deepspeed/inference/config.py @@ -310,6 +310,12 @@ def validate_dtype(cls, field_value, values): return field_value raise TypeError(f"Invalid type for dtype: {type(field_value)}") + @field_validator("max_out_tokens") + def validate_max_out_tokens(cls, field_value, values): + if field_value <= 0: + raise ValueError(f"max_out_tokens must be a positive integer, got {field_value}") + return field_value + @field_validator("moe") def moe_backward_compat(cls, field_value, values): if isinstance(field_value, bool): diff --git a/deepspeed/ops/transformer/inference/__init__.py b/deepspeed/ops/transformer/inference/__init__.py index c8b31a90eac2..ff5195e23ecf 100644 --- a/deepspeed/ops/transformer/inference/__init__.py +++ b/deepspeed/ops/transformer/inference/__init__.py @@ -4,5 +4,23 @@ # DeepSpeed Team from .config import DeepSpeedInferenceConfig -from ....model_implementations.transformers.ds_transformer import DeepSpeedTransformerInference from .moe_inference import DeepSpeedMoEInferenceConfig, DeepSpeedMoEInference + +__all__ = [ + "DeepSpeedInferenceConfig", + "DeepSpeedMoEInferenceConfig", + "DeepSpeedMoEInference", + "DeepSpeedTransformerInference", +] + + +def __getattr__(name: str): + # Lazy import breaks the circular dependency between this package and + # `deepspeed.model_implementations.transformers.ds_transformer`, which + # imports `deepspeed.ops.transformer.inference.triton.*` at module load + # time when Triton is installed. Accessing the symbol on demand keeps the + # package import path acyclic. + if name == "DeepSpeedTransformerInference": + from ....model_implementations.transformers.ds_transformer import DeepSpeedTransformerInference + return DeepSpeedTransformerInference + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/tests/unit/inference/test_inference_config.py b/tests/unit/inference/test_inference_config.py index fcba7e0d601b..4e4a7d2ad664 100644 --- a/tests/unit/inference/test_inference_config.py +++ b/tests/unit/inference/test_inference_config.py @@ -82,3 +82,27 @@ def test_moe_backward_compat_bool(self): config = DeepSpeedInferenceConfig(moe=value) assert isinstance(config.moe, DeepSpeedMoEConfig) assert config.moe.enabled == value + + +@pytest.mark.inference +class TestInferenceConfigValidation: + """CPU-only validation tests for DeepSpeedInferenceConfig.""" + + def test_negative_max_out_tokens_rejected(self): + # Regression test for https://github.com/deepspeedai/DeepSpeed/issues/8339 + from deepspeed.inference.config import DeepSpeedInferenceConfig + + with pytest.raises(ValueError, match="max_out_tokens must be a positive integer"): + DeepSpeedInferenceConfig(max_out_tokens=-1) + + def test_zero_max_out_tokens_rejected(self): + from deepspeed.inference.config import DeepSpeedInferenceConfig + + with pytest.raises(ValueError, match="max_out_tokens must be a positive integer"): + DeepSpeedInferenceConfig(max_out_tokens=0) + + def test_positive_max_out_tokens_accepted(self): + from deepspeed.inference.config import DeepSpeedInferenceConfig + + config = DeepSpeedInferenceConfig(max_out_tokens=128) + assert config.max_out_tokens == 128 diff --git a/tests/unit/inference/test_inference_import.py b/tests/unit/inference/test_inference_import.py new file mode 100644 index 000000000000..720056106d40 --- /dev/null +++ b/tests/unit/inference/test_inference_import.py @@ -0,0 +1,22 @@ +# Copyright (c) Microsoft Corporation. +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team + +import pytest + + +@pytest.mark.inference +class TestInferenceImport: + """CPU-only tests for the public inference package import surface.""" + + def test_transformer_inference_is_lazy_imported(self): + # Regression test for https://github.com/deepspeedai/DeepSpeed/issues/7159 + # Importing the inference package must not trigger a circular import even + # when Triton is installed, and the legacy public symbol must stay reachable. + from deepspeed.model_implementations.transformers.ds_transformer import ( + DeepSpeedTransformerInference as DirectTransformerInference, ) + from deepspeed.ops.transformer.inference import ( + DeepSpeedTransformerInference as OpsTransformerInference, ) + + assert OpsTransformerInference is DirectTransformerInference