diff --git a/src/lightmem/configs/pre_compressor/base.py b/src/lightmem/configs/pre_compressor/base.py index 523bc3a..d85fcbd 100644 --- a/src/lightmem/configs/pre_compressor/base.py +++ b/src/lightmem/configs/pre_compressor/base.py @@ -10,7 +10,7 @@ class PreCompressorConfig(BaseModel): _model_configs: ClassVar[Dict[str, str]] = { "llmlingua-2": "lightmem.configs.pre_compressor.llmlingua_2.LlmLingua2Config", - "entropy_compress": "lightmem.configs.pre_compressor.entropy_compress.EntropyCompressor" + "entropy_compress": "lightmem.configs.pre_compressor.entropy_compress.EntropyCompressorConfig" } configs: Dict[str, Any] = Field( @@ -47,4 +47,4 @@ def load_config_class(self) -> 'PreCompressorConfig': except (ImportError, AttributeError) as e: raise ValueError(f"Could not load config class '{config_path}': {e}") - return self \ No newline at end of file + return self diff --git a/src/lightmem/factory/pre_compressor/factory.py b/src/lightmem/factory/pre_compressor/factory.py index 8e4fef6..3eb8da5 100644 --- a/src/lightmem/factory/pre_compressor/factory.py +++ b/src/lightmem/factory/pre_compressor/factory.py @@ -5,7 +5,7 @@ class PreCompressorFactory: _MODEL_MAPPING: Dict[str, str] = { "llmlingua-2": "lightmem.factory.pre_compressor.llmlingua_2.LlmLingua2Compressor", - "entropy_compress": "lightmem.factory.pre_compressor.entropy_compress.", + "entropy_compress": "lightmem.factory.pre_compressor.entropy_compress.EntropyCompressor", } @classmethod @@ -53,4 +53,4 @@ def from_config(cls, config: PreCompressorConfig): except Exception as e: raise ValueError( f"Failed to instantiate {model_name} compressor: {str(e)}" - ) from e \ No newline at end of file + ) from e diff --git a/tests/test_pre_compressor_factory.py b/tests/test_pre_compressor_factory.py new file mode 100644 index 0000000..1ca99d7 --- /dev/null +++ b/tests/test_pre_compressor_factory.py @@ -0,0 +1,37 @@ +from types import SimpleNamespace + +from lightmem.configs.pre_compressor.base import PreCompressorConfig +from lightmem.configs.pre_compressor.entropy_compress import EntropyCompressorConfig +from lightmem.factory.pre_compressor.factory import PreCompressorFactory + + +def test_entropy_compressor_config_resolves_to_its_config_class(): + config = PreCompressorConfig(model_name="entropy_compress") + + assert isinstance(config.configs, EntropyCompressorConfig) + + +def test_entropy_compressor_factory_resolves_to_its_implementation(monkeypatch): + compressor_config = object() + imported = {} + + class FakeEntropyCompressor: + def __init__(self, config=None): + self.config = config + + def fake_import_module(module_path): + imported["module_path"] = module_path + return SimpleNamespace(EntropyCompressor=FakeEntropyCompressor) + + monkeypatch.setattr( + "lightmem.factory.pre_compressor.factory.import_module", + fake_import_module, + ) + + compressor = PreCompressorFactory.from_config( + SimpleNamespace(model_name="entropy_compress", configs=compressor_config) + ) + + assert imported["module_path"] == "lightmem.factory.pre_compressor.entropy_compress" + assert isinstance(compressor, FakeEntropyCompressor) + assert compressor.config is compressor_config