Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
85 changes: 85 additions & 0 deletions mergekit/_data/architectures/granite.json
Original file line number Diff line number Diff line change
@@ -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"
]
}
]
}
5 changes: 5 additions & 0 deletions mergekit/architecture/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()
16 changes: 16 additions & 0 deletions tests/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
CLIPVisionConfig,
GPT2Config,
GPT2LMHeadModel,
GraniteConfig,
GraniteForCausalLM,
LlamaConfig,
LlamaForCausalLM,
LlavaConfig,
Expand Down Expand Up @@ -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,
Expand Down
45 changes: 45 additions & 0 deletions tests/test_basic_merges.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Loading