Skip to content
Open
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
5 changes: 4 additions & 1 deletion litgpt/finetune/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
from litgpt.constants import _BITANDBYTES_AVAILABLE_NOT_EQUAL_0_42_0
from litgpt.data import Alpaca, DataModule
from litgpt.generate.base import generate
from litgpt.lora import GPT, Block, Config, mark_only_lora_as_trainable
from litgpt.lora import GPT, Block, Config, mark_only_lora_as_trainable, normalize_lora_r_by_layer
from litgpt.parser_config import save_hyperparameters
from litgpt.prompts import save_prompt_style
from litgpt.scripts.merge_lora import merge_lora
Expand Down Expand Up @@ -54,6 +54,7 @@ def setup(
devices: int | str = 1,
num_nodes: int = 1,
lora_r: int = 8,
lora_r_by_layer: dict[int, int] | None = None,
lora_alpha: int = 16,
lora_dropout: float = 0.05,
lora_query: bool = True,
Expand Down Expand Up @@ -91,6 +92,7 @@ def setup(
devices: How many devices/GPUs to use.
num_nodes: How many nodes the code is being run on.
lora_r: The LoRA rank.
lora_r_by_layer: The LoRA rank for each layer (zero-based index). lora_r is used for non-specified layers.
lora_alpha: The LoRA alpha.
lora_dropout: The LoRA dropout value.
lora_query: Whether to apply LoRA to the query weights in attention.
Expand Down Expand Up @@ -118,6 +120,7 @@ def setup(
config = Config.from_file(
checkpoint_dir / "model_config.yaml",
lora_r=lora_r,
lora_r_by_layer=normalize_lora_r_by_layer(lora_r_by_layer),
lora_alpha=lora_alpha,
lora_dropout=lora_dropout,
lora_query=lora_query,
Expand Down
102 changes: 89 additions & 13 deletions litgpt/lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -451,6 +451,34 @@ def lora_filter(key: str, value: Any) -> bool:
return "lora_" in key


def normalize_lora_r_by_layer(
overrides: dict[int | str, int] | None,
) -> dict[int, int] | None:
"""Normalize serialized layer indices to integers."""
if overrides is None:
return None

if not isinstance(overrides, dict):
raise TypeError("lora_r_by_layer must be a dictionary")

lora_r_by_layer = {}

for key, rank in overrides.items():
if type(key) is int:
block_idx = key
elif isinstance(key, str):
try:
block_idx = int(key)
except ValueError:
raise TypeError(f"Invalid LoRA block index: {key}") from None
else:
raise TypeError(f"Invalid LoRA block index: {key}")

lora_r_by_layer[block_idx] = rank

return lora_r_by_layer


@dataclass
class Config(BaseConfig):
"""
Expand All @@ -467,13 +495,58 @@ class Config(BaseConfig):
lora_r: int = 0
lora_alpha: int = 1
lora_dropout: float = 0.0

# optional rank overrides for individual transformer blocks
lora_r_by_layer: dict[int, int] | None = None

lora_query: bool = False
lora_key: bool = False
lora_value: bool = False
lora_projection: bool = False
lora_mlp: bool = False
lora_head: bool = False

def __post_init__(self) -> None:
super().__post_init__()

# validate the global rank
if type(self.lora_r) is not int:
raise TypeError("lora_r must be an integer")

if self.lora_r < 0:
raise ValueError("lora_r must be non-negative")

# presrve the default lora definition when lora_r_by_layer is not set
if self.lora_r_by_layer is None:
return

if not isinstance(self.lora_r_by_layer, dict):
raise TypeError("lora_r_by_layer must be a dictionary")

for block_idx, rank in self.lora_r_by_layer.items():
if type(block_idx) is not int:
raise TypeError(f"LoRA block index must be an integer, got {block_idx}")

if not 0 <= block_idx < self.n_layer:
raise ValueError(f"LoRA block index {block_idx} is out of range for a model with {self.n_layer} layers")

if type(rank) is not int:
raise TypeError(f"LoRA rank for block {block_idx} must be an integer")

if rank < 0:
raise ValueError(f"LoRA rank for block {block_idx} must be non-negative")

def get_lora_rank(self, block_idx: int | None = None) -> int:
"""
Return the LoRA rank for a given transformer block.

Fall back to the global rank when no override exists or when the module does not belong to a transformer block.
"""
if block_idx is None or self.lora_r_by_layer is None:
return self.lora_r

return self.lora_r_by_layer.get(block_idx, self.lora_r)

@property
def mlp_class(self) -> type:
return getattr(litgpt.lora, self.mlp_class_name)
Expand Down Expand Up @@ -524,7 +597,7 @@ class Block(BaseBlock):
def __init__(self, config: Config, block_idx: int) -> None:
super().__init__(config, block_idx)
self.attn = CausalSelfAttention(config, block_idx)
self.mlp = config.mlp_class(config)
self.mlp = config.mlp_class(config, block_idx=block_idx)


class CausalSelfAttention(BaseCausalSelfAttention):
Expand All @@ -535,7 +608,7 @@ def __init__(self, config: Config, block_idx: int) -> None:
self.qkv = LoRAQKVLinear(
in_features=config.n_embd,
out_features=shape,
r=config.lora_r,
r=config.get_lora_rank(block_idx),
lora_alpha=config.lora_alpha,
lora_dropout=config.lora_dropout,
enable_lora=(config.lora_query, config.lora_key, config.lora_value),
Expand All @@ -551,6 +624,7 @@ def __init__(self, config: Config, block_idx: int) -> None:
config.head_size * config.n_head,
config.n_embd,
use_r=config.lora_projection,
block_idx=block_idx,
)

def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None:
Expand Down Expand Up @@ -578,6 +652,7 @@ def create_lora_linear(
out_size: int,
bias: float | bool | None = None,
use_r: bool | None = None,
block_idx: int | None = None,
) -> LoRALinear:
if bias is None:
bias = config.bias
Expand All @@ -587,17 +662,17 @@ def create_lora_linear(
in_size,
out_size,
bias=bias,
r=(config.lora_r if use_r else 0),
r=(config.get_lora_rank(block_idx) if use_r else 0),
lora_alpha=config.lora_alpha,
lora_dropout=config.lora_dropout,
)


class GptNeoxMLP(litgpt.model.GptNeoxMLP):
def __init__(self, config: Config) -> None:
def __init__(self, config: Config, block_idx: int | None = None) -> None:
nn.Module.__init__(self)
self.fc = create_lora_linear(config, config.n_embd, config.intermediate_size)
self.proj = create_lora_linear(config, config.intermediate_size, config.n_embd)
self.fc = create_lora_linear(config, config.n_embd, config.intermediate_size, block_idx=block_idx)
self.proj = create_lora_linear(config, config.intermediate_size, config.n_embd, block_idx=block_idx)
self.config = config

def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None:
Expand All @@ -613,12 +688,12 @@ def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwa


class LLaMAMLP(litgpt.model.LLaMAMLP):
def __init__(self, config: Config, intermediate_size: int | None = None) -> None:
def __init__(self, config: Config, intermediate_size: int | None = None, block_idx: int | None = None) -> None:
nn.Module.__init__(self)
self.intermediate_size = intermediate_size or config.intermediate_size
self.fc_1 = create_lora_linear(config, config.n_embd, self.intermediate_size)
self.fc_2 = create_lora_linear(config, config.n_embd, self.intermediate_size)
self.proj = create_lora_linear(config, self.intermediate_size, config.n_embd)
self.fc_1 = create_lora_linear(config, config.n_embd, self.intermediate_size, block_idx=block_idx)
self.fc_2 = create_lora_linear(config, config.n_embd, self.intermediate_size, block_idx=block_idx)
self.proj = create_lora_linear(config, self.intermediate_size, config.n_embd, block_idx=block_idx)
self.config = config

def _load_from_state_dict(self, state_dict: dict, prefix: str, *args: Any, **kwargs: Any) -> None:
Expand All @@ -644,11 +719,12 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:


class LLaMAMoE(litgpt.model.LLaMAMoE):
def __init__(self, config: Config) -> None:
def __init__(self, config: Config, block_idx: int | None = None) -> None:
nn.Module.__init__(self)
self.gate = create_lora_linear(config, config.n_embd, config.n_expert, bias=False)
self.gate = create_lora_linear(config, config.n_embd, config.n_expert, bias=False, block_idx=block_idx)
self.experts = nn.ModuleList(
LLaMAMLP(config, intermediate_size=config.moe_intermediate_size) for _ in range(config.n_expert)
LLaMAMLP(config, intermediate_size=config.moe_intermediate_size, block_idx=block_idx)
for _ in range(config.n_expert)
)
self.config = config

Expand Down
4 changes: 3 additions & 1 deletion litgpt/scripts/merge_lora.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
import torch
import yaml

from litgpt.lora import GPT, Config, lora_filter, merge_lora_weights
from litgpt.lora import GPT, Config, lora_filter, merge_lora_weights, normalize_lora_r_by_layer
from litgpt.utils import check_valid_checkpoint_dir, extend_checkpoint_dir


Expand Down Expand Up @@ -96,6 +96,8 @@ def load_lora_metadata(checkpoint_dir: Path) -> tuple[dict[str, Any], Path, str
hparams = yaml.safe_load(file)

lora_params = {k: v for k, v in hparams.items() if k.startswith("lora_")}
if "lora_r_by_layer" in lora_params:
lora_params["lora_r_by_layer"] = normalize_lora_r_by_layer(lora_params["lora_r_by_layer"])
pretrained_checkpoint_dir = Path(hparams["checkpoint_dir"])
precision = hparams.get("precision")
return lora_params, pretrained_checkpoint_dir, precision
Loading