diff --git a/tunix/models/gemma4/__init__.py b/tunix/models/gemma4/__init__.py index cd9687b35..3497018ff 100644 --- a/tunix/models/gemma4/__init__.py +++ b/tunix/models/gemma4/__init__.py @@ -15,11 +15,9 @@ """Gemma4 API.""" from tunix.models.gemma4 import mapping_vllm_jax -from tunix.models.gemma4 import model -from tunix.models.gemma4 import params_safetensors BACKEND_MAPPINGS = { 'vllm_jax': mapping_vllm_jax.VLLM_JAX_MAPPING, } -__all__ = ['BACKEND_MAPPINGS', 'model', 'params_safetensors'] +__all__ = ['BACKEND_MAPPINGS'] diff --git a/tunix/models/gemma4/model.py b/tunix/models/gemma4/model.py index 69ca6d083..5a19ff2e5 100644 --- a/tunix/models/gemma4/model.py +++ b/tunix/models/gemma4/model.py @@ -269,6 +269,13 @@ def gemma4_e4b( audio_encoder=audio.ConformerConfig(), ) + @classmethod + def gemma4_e4b_it( + cls, + sharding_config: ShardingConfig = ShardingConfig.get_default_sharding(), + ) -> 'ModelConfig': + return cls.gemma4_e4b(sharding_config=sharding_config) + @classmethod def gemma4_12b( cls, @@ -1538,6 +1545,7 @@ def init_cache(self, batch_size, max_seq_len, dtype): class Gemma4(BackendMappingMixin, nnx.Module): """Gemma4 model.""" + BACKEND_PACKAGE_PATH = __name__ def __init__( self, config: ModelConfig, *, rngs: nnx.Rngs, text_only: bool = True @@ -1873,4 +1881,4 @@ def get_model_input(self): @property def num_embed(self) -> int: - return self.config.num_embed + return self.config.num_embed \ No newline at end of file