| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from transformers.models.gemma4.modeling_gemma4 import ( |
| Gemma4TextDecoderLayer, |
| Gemma4ForCausalLM, |
| Gemma4TextModel, |
| ) |
| from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig |
|
|
|
|
| class SolonMoEConfig(Gemma4TextConfig): |
| model_type = "solon_moe" |
| |
|
|
|
|
| class SolonMoEDecoderLayer(Gemma4TextDecoderLayer): |
| def __init__(self, config, layer_idx): |
| moe_layers = getattr(config, "moe_layers", None) |
| want = (moe_layers is None and config.enable_moe_block) or \ |
| (moe_layers is not None and layer_idx in set(moe_layers)) |
| saved = config.enable_moe_block |
| config.enable_moe_block = bool(want) |
| try: |
| super().__init__(config, layer_idx) |
| finally: |
| config.enable_moe_block = saved |
| self.enable_moe_block = bool(want) |
|
|
|
|
| class SolonMoETextModel(Gemma4TextModel): |
| config_class = SolonMoEConfig |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| |
| import torch.nn as nn |
| self.layers = nn.ModuleList( |
| [SolonMoEDecoderLayer(config, i) for i in range(config.num_hidden_layers)] |
| ) |
| self.post_init() |
|
|
|
|
| class SolonMoEForCausalLM(Gemma4ForCausalLM): |
| config_class = SolonMoEConfig |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.model = SolonMoETextModel(config) |
| self.post_init() |
|
|