Ouzhang's picture
Add files using upload-large-folder tool
31dc8dc verified
Raw
History Blame Contribute Delete
10.5 kB
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
from diffulex.distributed.parallel_state import fetch_parallel_state
from diffulex.vllm_compat import get_vllm_tp_group
def tp_all_reduce(x: torch.Tensor, group) -> torch.Tensor:
vllm_tp_group = get_vllm_tp_group()
if vllm_tp_group is not None:
return vllm_tp_group.all_reduce(x)
dist.all_reduce(x, group=group)
return x
def divide(numerator, denominator):
assert numerator % denominator == 0
return numerator // denominator
class LoRAMixin:
"""Mixin class to add LoRA support to existing linear layers."""
def __init_lora__(self, r: int = 0, lora_alpha: float = 1.0, lora_dropout: float = 0.0):
if r > 0:
self.r = r
self.lora_alpha = lora_alpha
self.scaling = lora_alpha / r
# Initialize LoRA parameters
if hasattr(self, "output_size_per_partition"):
out_features = self.output_size_per_partition
else:
out_features = self.output_size
if hasattr(self, "input_size_per_partition"):
in_features = self.input_size_per_partition
else:
in_features = self.input_size
self.lora_A = nn.Parameter(torch.zeros(r, in_features))
self.lora_B = nn.Parameter(torch.zeros(out_features, r))
self.lora_dropout = nn.Dropout(lora_dropout) if lora_dropout > 0 else nn.Identity()
self.merged = False
# Initialize weights
nn.init.kaiming_uniform_(self.lora_A, a=5**0.5)
nn.init.zeros_(self.lora_B)
else:
self.r = 0
self.merged = True
def merge_lora(self):
"""Merge LoRA weights into base weight."""
if not (hasattr(self, "r") and self.r > 0 and not self.merged):
return
# If base weight is missing, we cannot merge in-place. Keep LoRA unmerged and apply via lora_forward.
weight = getattr(self, "weight", None)
if weight is None or not hasattr(weight, "data"):
return
self.weight.data += self.scaling * torch.mm(self.lora_B, self.lora_A)
self.merged = True
def lora_forward(self, x: torch.Tensor, base_output: torch.Tensor) -> torch.Tensor:
"""Apply LoRA forward pass."""
if not hasattr(self, "r") or self.r == 0 or self.merged:
return base_output
lora_out = F.linear(self.lora_dropout(x), self.lora_A)
lora_out = F.linear(lora_out, self.lora_B)
return base_output + lora_out * self.scaling
class LinearBase(nn.Module):
def __init__(
self,
input_size: int,
output_size: int,
tp_dim: int | None = None,
):
super().__init__()
self.input_size = input_size
self.output_size = output_size
self.tp_dim = tp_dim
parallel_state = fetch_parallel_state()
self.tp_rank = parallel_state.get_tp_rank()
self.tp_size = parallel_state.get_tp_world_size()
self.tp_group = parallel_state.get_tp_group()
def _forward_base(self, x: torch.Tensor, bias: nn.Parameter | None) -> torch.Tensor:
return F.linear(x, self.weight, bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
raise NotImplementedError
class ReplicatedLinear(LinearBase, LoRAMixin):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = False,
r: int = 0,
lora_alpha: float = 1.0,
lora_dropout: float = 0.0,
):
LinearBase.__init__(self, input_size, output_size, None)
self.weight = nn.Parameter(torch.empty(self.output_size, self.input_size))
self.weight.weight_loader = self.weight_loader
if bias:
self.bias = nn.Parameter(torch.empty(self.output_size))
self.bias.weight_loader = self.weight_loader
else:
self.register_parameter("bias", None)
self.__init_lora__(r, lora_alpha, lora_dropout)
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
param.data.copy_(loaded_weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
base_out = self._forward_base(x, self.bias)
return self.lora_forward(x, base_out)
class ColumnParallelLinear(LinearBase, LoRAMixin):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = False,
r: int = 0,
lora_alpha: float = 1.0,
lora_dropout: float = 0.0,
):
LinearBase.__init__(self, input_size, output_size, 0)
self.input_size_per_partition = input_size
self.output_size_per_partition = divide(output_size, self.tp_size)
self._forward_out_features = int(self.output_size_per_partition)
self.weight = nn.Parameter(torch.empty(self.output_size_per_partition, self.input_size))
self.weight.weight_loader = self.weight_loader
if bias:
self.bias = nn.Parameter(torch.empty(self.output_size_per_partition))
self.bias.weight_loader = self.weight_loader
else:
self.register_parameter("bias", None)
self.__init_lora__(r, lora_alpha, lora_dropout)
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
param_data = param.data
shard_size = param_data.size(self.tp_dim)
start_idx = self.tp_rank * shard_size
loaded_weight = loaded_weight.narrow(self.tp_dim, start_idx, shard_size)
param_data.copy_(loaded_weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
base_out = self._forward_base(x, self.bias)
return self.lora_forward(x, base_out)
class MergedColumnParallelLinear(ColumnParallelLinear):
def __init__(
self,
input_size: int,
output_sizes: list[int],
bias: bool = False,
r: int = 0,
lora_alpha: float = 1.0,
lora_dropout: float = 0.0,
):
self.output_sizes = output_sizes
super().__init__(
input_size,
sum(output_sizes),
bias=bias,
r=r,
lora_alpha=lora_alpha,
lora_dropout=lora_dropout,
)
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: int):
param_data = param.data
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // self.tp_size
shard_size = self.output_sizes[loaded_shard_id] // self.tp_size
param_data = param_data.narrow(self.tp_dim, shard_offset, shard_size)
loaded_weight = loaded_weight.chunk(self.tp_size, self.tp_dim)[self.tp_rank]
param_data.copy_(loaded_weight)
class QKVParallelLinear(ColumnParallelLinear):
def __init__(
self,
hidden_size: int,
head_size: int,
total_num_heads: int,
total_num_kv_heads: int | None = None,
bias: bool = False,
r: int = 0,
lora_alpha: float = 1.0,
lora_dropout: float = 0.0,
):
self.head_size = head_size
self.total_num_heads = total_num_heads
self.total_num_kv_heads = total_num_kv_heads or total_num_heads
parallel_state = fetch_parallel_state()
tp_size = parallel_state.get_tp_world_size()
self.num_heads = divide(self.total_num_heads, tp_size)
self.num_kv_heads = divide(self.total_num_kv_heads, tp_size)
input_size = hidden_size
output_size = (self.total_num_heads + 2 * self.total_num_kv_heads) * self.head_size
super().__init__(
input_size,
output_size,
bias,
r,
lora_alpha,
lora_dropout,
)
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: str):
param_data = param.data
assert loaded_shard_id in ["q", "k", "v"]
if loaded_shard_id == "q":
shard_size = self.num_heads * self.head_size
shard_offset = 0
elif loaded_shard_id == "k":
shard_size = self.num_kv_heads * self.head_size
shard_offset = self.num_heads * self.head_size
else:
shard_size = self.num_kv_heads * self.head_size
shard_offset = self.num_heads * self.head_size + self.num_kv_heads * self.head_size
param_data = param_data.narrow(self.tp_dim, shard_offset, shard_size)
loaded_weight = loaded_weight.chunk(self.tp_size, self.tp_dim)[self.tp_rank]
param_data.copy_(loaded_weight)
class RowParallelLinear(LinearBase, LoRAMixin):
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = False,
r: int = 0,
lora_alpha: float = 1.0,
lora_dropout: float = 0.0,
):
LinearBase.__init__(self, input_size, output_size, 1)
self.input_size_per_partition = divide(input_size, self.tp_size)
self.output_size_per_partition = output_size
self.weight = nn.Parameter(torch.empty(self.output_size, self.input_size_per_partition))
self.weight.weight_loader = self.weight_loader
if bias:
self.bias = nn.Parameter(torch.empty(self.output_size))
self.bias.weight_loader = self.weight_loader
else:
self.register_parameter("bias", None)
self.__init_lora__(r, lora_alpha, lora_dropout)
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
param_data = param.data
shard_size = param_data.size(self.tp_dim)
start_idx = self.tp_rank * shard_size
loaded_weight = loaded_weight.narrow(self.tp_dim, start_idx, shard_size)
param_data.copy_(loaded_weight)
def forward(self, x: torch.Tensor) -> torch.Tensor:
bias = self.bias if self.tp_rank == 0 else None
tp_group = self.tp_group
y = self._forward_base(x, bias)
if hasattr(self, "r") and self.r > 0 and not self.merged:
lora_out = F.linear(self.lora_dropout(x), self.lora_A)
lora_out = F.linear(lora_out, self.lora_B)
if self.tp_size > 1:
y = tp_all_reduce(y, tp_group)
lora_out = tp_all_reduce(lora_out, tp_group)
return y + lora_out * self.scaling
if self.tp_size > 1:
y = tp_all_reduce(y, tp_group)
return y