diff --git a/mergekit/_data/architectures/granite.json b/mergekit/_data/architectures/granite.json new file mode 100644 index 00000000..a9e658ad --- /dev/null +++ b/mergekit/_data/architectures/granite.json @@ -0,0 +1,85 @@ +{ + "model_type": "granite", + "architectures": [ + "GraniteForCausalLM" + ], + "pre_weights": [ + { + "name": "model.embed_tokens.weight", + "is_embed": true + } + ], + "num_layers_config_key": "num_hidden_layers", + "layer_templates": { + "weights": [ + { + "name": "model.layers.${layer_index}.input_layernorm.weight" + }, + { + "name": "model.layers.${layer_index}.self_attn.q_proj.weight" + }, + { + "name": "model.layers.${layer_index}.self_attn.q_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.self_attn.k_proj.weight" + }, + { + "name": "model.layers.${layer_index}.self_attn.k_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.self_attn.v_proj.weight" + }, + { + "name": "model.layers.${layer_index}.self_attn.v_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.self_attn.o_proj.weight" + }, + { + "name": "model.layers.${layer_index}.self_attn.o_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.post_attention_layernorm.weight" + }, + { + "name": "model.layers.${layer_index}.mlp.gate_proj.weight" + }, + { + "name": "model.layers.${layer_index}.mlp.gate_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.mlp.up_proj.weight" + }, + { + "name": "model.layers.${layer_index}.mlp.up_proj.bias", + "optional": true + }, + { + "name": "model.layers.${layer_index}.mlp.down_proj.weight" + }, + { + "name": "model.layers.${layer_index}.mlp.down_proj.bias", + "optional": true + } + ] + }, + "post_weights": [ + { + "name": "model.norm.weight" + }, + { + "name": "lm_head.weight", + "is_embed": true, + "optional": true, + "tied_names": [ + "model.embed_tokens.weight" + ] + } + ] +} diff --git a/mergekit/architecture/base.py b/mergekit/architecture/base.py index 5fca5bdb..738fb8f6 100644 --- a/mergekit/architecture/base.py +++ b/mergekit/architecture/base.py @@ -4,6 +4,7 @@ from abc import ABC, abstractmethod from typing import Dict, List, Optional, Tuple +import torch # required for Pydantic to resolve PretrainedConfig's torch.dtype forward reference from pydantic import BaseModel, Field from transformers import PretrainedConfig @@ -151,3 +152,7 @@ def get_module(self, module_name: str) -> ConfiguredModuleArchitecture: config=self.config, weight_prefix=self.info.modules[module_name].weight_prefix, ) + + +ConfiguredModuleArchitecture.model_rebuild() +ConfiguredModelArchitecture.model_rebuild() diff --git a/tests/common.py b/tests/common.py index 9b7ceb9c..13cb951d 100644 --- a/tests/common.py +++ b/tests/common.py @@ -7,6 +7,8 @@ CLIPVisionConfig, GPT2Config, GPT2LMHeadModel, + GraniteConfig, + GraniteForCausalLM, LlamaConfig, LlamaForCausalLM, LlavaConfig, @@ -86,6 +88,20 @@ def make_picollama(path: str, vocab_size: int = 64): return str(path) +def make_picogranite(path: str, vocab_size: int = 64): + cfg = GraniteConfig( + vocab_size=vocab_size, + hidden_size=32, + intermediate_size=48, + num_attention_heads=2, + num_key_value_heads=2, + num_hidden_layers=2, + ) + model = GraniteForCausalLM(cfg) + model.save_pretrained(path, safe_serialization=True) + return str(path) + + def make_gpt2size(path: str): cfg = GPT2Config( n_ctx=1024, diff --git a/tests/test_basic_merges.py b/tests/test_basic_merges.py index ab37d981..f310b409 100644 --- a/tests/test_basic_merges.py +++ b/tests/test_basic_merges.py @@ -13,6 +13,7 @@ from mergekit.io import LazyTensorLoader from tests.common import ( make_gpt2size, + make_picogranite, make_picollama, make_picoLlaVa, run_and_check_merge, @@ -54,6 +55,50 @@ def gpt2_like(tmp_path_factory): return make_gpt2size(tmp_path_factory.mktemp("gpt2_like")) +@pytest.fixture(scope="session") +def granite_a(tmp_path_factory): + return make_picogranite(tmp_path_factory.mktemp("granite_a")) + + +@pytest.fixture(scope="session") +def granite_b(tmp_path_factory): + return make_picogranite(tmp_path_factory.mktemp("granite_b")) + + +class TestGraniteMerges: + def test_granite_copy(self, granite_a): + config = MergeConfiguration( + merge_method="passthrough", + models=[InputModelDefinition(model=granite_a)], + dtype="bfloat16", + ) + run_and_check_merge(config) + + def test_granite_linear_merge(self, granite_a, granite_b): + config = MergeConfiguration( + merge_method="linear", + models=[ + InputModelDefinition(model=granite_a, parameters={"weight": 0.6}), + InputModelDefinition(model=granite_b, parameters={"weight": 0.4}), + ], + dtype="bfloat16", + ) + run_and_check_merge(config) + + def test_granite_slerp(self, granite_a, granite_b): + config = MergeConfiguration( + merge_method="slerp", + base_model=granite_a, + models=[ + InputModelDefinition(model=granite_a), + InputModelDefinition(model=granite_b), + ], + parameters={"t": 0.5}, + dtype="bfloat16", + ) + run_and_check_merge(config) + + class TestBasicMerges: def test_gpt2_copy(self, gpt2_like): config = MergeConfiguration(