[lora] More efficient pinned memory (#20876)
This commit is contained in:
@@ -48,6 +48,7 @@ class LoRALayer(nn.Module):
|
||||
|
||||
# lora weights in cpu. The weights are loaded from checkpoint.
|
||||
self.weights: Dict[str, torch.Tensor] = {}
|
||||
self.pinned_weights: Dict[str, torch.Tensor] = {}
|
||||
|
||||
|
||||
class LoRAAdapter(nn.Module):
|
||||
@@ -87,7 +88,9 @@ class LoRAAdapter(nn.Module):
|
||||
)
|
||||
|
||||
self.embedding_layers: Dict[str, torch.Tensor] = {}
|
||||
self.pinned_embedding_layers: Dict[str, torch.Tensor] = {}
|
||||
self.added_tokens_embeddings: Dict[str, torch.Tensor] = {}
|
||||
self.pinned_added_tokens_embeddings: Dict[str, torch.Tensor] = {}
|
||||
|
||||
@staticmethod
|
||||
def _build_moe_gated_map(base_model: torch.nn.Module) -> Dict[int, bool]:
|
||||
|
||||
@@ -631,10 +631,6 @@ class LoRAManager:
|
||||
)
|
||||
lora_adapter.initialize_weights()
|
||||
|
||||
# If we want to overlap loading LoRA adapters with compute, they must be pinned in CPU memory
|
||||
if self.enable_lora_overlap_loading:
|
||||
lora_adapter.pin_weights_in_cpu()
|
||||
|
||||
self.loras[lora_ref.lora_id] = lora_adapter
|
||||
|
||||
def load_lora_weights_from_tensors(
|
||||
@@ -707,6 +703,7 @@ class LoRAManager:
|
||||
lora_added_tokens_size=self.lora_added_tokens_size,
|
||||
experts_shared_outer_loras=self.experts_shared_outer_loras,
|
||||
strict_loading=self.lora_strict_loading,
|
||||
enable_lora_overlap_loading=self.enable_lora_overlap_loading,
|
||||
)
|
||||
|
||||
# Initializing memory pool with base model
|
||||
|
||||
@@ -1,6 +1,17 @@
|
||||
import logging
|
||||
import re
|
||||
from typing import Callable, Dict, Iterable, Iterator, List, Optional, Set, Tuple, Union
|
||||
from typing import (
|
||||
Callable,
|
||||
Dict,
|
||||
Iterable,
|
||||
Iterator,
|
||||
List,
|
||||
Optional,
|
||||
Set,
|
||||
Tuple,
|
||||
Union,
|
||||
overload,
|
||||
)
|
||||
|
||||
import torch
|
||||
|
||||
@@ -22,12 +33,14 @@ from sglang.srt.lora.utils import (
|
||||
REPLICATED_LINEAR_LORA_NAMES,
|
||||
ROW_PARALLELISM_LINEAR_LORA_NAMES,
|
||||
LoRAType,
|
||||
copy_weight_into_buffer,
|
||||
get_hidden_dim,
|
||||
get_lm_head_lora_b_shard_size,
|
||||
get_normalized_target_modules,
|
||||
get_stacked_multiply,
|
||||
get_target_module_name,
|
||||
)
|
||||
from sglang.srt.utils import is_pin_memory_available
|
||||
from sglang.srt.utils.hf_transformers_utils import AutoConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -53,6 +66,28 @@ class EmptySlot:
|
||||
EMPTY_SLOT = EmptySlot()
|
||||
|
||||
|
||||
@overload
|
||||
def append_cache_key_suffix(cache_keys: str, suffix: str) -> str: ...
|
||||
|
||||
|
||||
@overload
|
||||
def append_cache_key_suffix(
|
||||
cache_keys: Dict[int, str], suffix: str
|
||||
) -> Dict[int, str]: ...
|
||||
|
||||
|
||||
def append_cache_key_suffix(
|
||||
cache_keys: Union[str, Dict[int, str]],
|
||||
suffix: str,
|
||||
) -> Union[str, Dict[int, str]]:
|
||||
if isinstance(cache_keys, dict):
|
||||
return {
|
||||
expert_id: f"{cache_key}#{suffix}"
|
||||
for expert_id, cache_key in cache_keys.items()
|
||||
}
|
||||
return f"{cache_keys}#{suffix}"
|
||||
|
||||
|
||||
def _get_moe_ep_context() -> Tuple[int, int]:
|
||||
"""Return `(moe_ep_size, moe_ep_rank)`, or `(1, 0)` if the MoE EP group
|
||||
is not initialized (hermetic tests or pure-TP launches)."""
|
||||
@@ -107,6 +142,7 @@ class LoRAMemoryPool:
|
||||
lora_added_tokens_size: int,
|
||||
experts_shared_outer_loras: bool = False,
|
||||
strict_loading: bool = False,
|
||||
enable_lora_overlap_loading: bool = False,
|
||||
):
|
||||
self.base_hf_config: AutoConfig = base_hf_config
|
||||
self.num_layer: int = base_hf_config.num_hidden_layers
|
||||
@@ -119,6 +155,8 @@ class LoRAMemoryPool:
|
||||
self.target_modules: Set[str] = target_modules
|
||||
self.experts_shared_outer_loras: bool = experts_shared_outer_loras
|
||||
self.strict_loading: bool = strict_loading
|
||||
self.enable_lora_overlap_loading: bool = enable_lora_overlap_loading
|
||||
self.pin_memory_available: bool = is_pin_memory_available()
|
||||
|
||||
# Under EP with a Triton/DeepGEMM runner, `StandardDispatcher` remaps
|
||||
# global `topk_ids` -> local expert IDs before the MoE kernel, so
|
||||
@@ -258,18 +296,21 @@ class LoRAMemoryPool:
|
||||
def _iter_local_expert_weights(
|
||||
self,
|
||||
weights: Union[torch.Tensor, Dict[int, torch.Tensor]],
|
||||
) -> Iterator[Tuple[int, torch.Tensor]]:
|
||||
"""Yield `(local_expert_id, weight)` pairs for per-expert MoE LoRA
|
||||
inputs, filtered/remapped to this rank's slice. Accepts either a
|
||||
`{global_eid: 2D tensor}` dict or a 3D `[num_experts, *, *]` tensor."""
|
||||
cache_keys: Union[str, Dict[int, str]],
|
||||
) -> Iterator[Tuple[int, torch.Tensor, str]]:
|
||||
"""Yield `(local_expert_id, weight, cache_key)` triples for per-expert
|
||||
MoE LoRA inputs, filtered/remapped to this rank's slice. Accepts either
|
||||
a `{global_eid: 2D tensor}` dict or a 3D `[num_experts, *, *]` tensor."""
|
||||
if isinstance(weights, dict):
|
||||
assert isinstance(cache_keys, dict)
|
||||
for gid, w in weights.items():
|
||||
lid = self._global_to_local_expert_id(gid)
|
||||
if lid is not None:
|
||||
yield lid, w
|
||||
yield lid, w, cache_keys[gid]
|
||||
return
|
||||
|
||||
if isinstance(weights, torch.Tensor) and weights.dim() == 3:
|
||||
assert isinstance(cache_keys, str)
|
||||
total = weights.shape[0]
|
||||
if self.moe_use_local_expert_ids:
|
||||
start = self.moe_ep_rank * self._num_experts_local
|
||||
@@ -277,7 +318,11 @@ class LoRAMemoryPool:
|
||||
else:
|
||||
start, count = 0, total
|
||||
for i in range(count):
|
||||
yield i, weights[start + i]
|
||||
yield (
|
||||
i,
|
||||
weights[start + i],
|
||||
append_cache_key_suffix(cache_keys, f"expert{start + i}"),
|
||||
)
|
||||
return
|
||||
|
||||
raise TypeError(
|
||||
@@ -593,6 +638,35 @@ class LoRAMemoryPool:
|
||||
self.get_lora_B_shape,
|
||||
)
|
||||
|
||||
def _get_maybe_cached_weight_for_transfer(
|
||||
self,
|
||||
pinned_weight_store: Dict[str, torch.Tensor],
|
||||
cache_key: str,
|
||||
weight: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
if (
|
||||
not self.pin_memory_available
|
||||
or weight.device.type != "cpu"
|
||||
or weight.is_pinned()
|
||||
):
|
||||
return weight
|
||||
|
||||
if not self.enable_lora_overlap_loading:
|
||||
return weight.pin_memory()
|
||||
|
||||
cached_weight = pinned_weight_store.get(cache_key)
|
||||
if cached_weight is None:
|
||||
cached_weight = weight.pin_memory()
|
||||
pinned_weight_store[cache_key] = cached_weight
|
||||
elif cached_weight.shape != weight.shape or cached_weight.dtype != weight.dtype:
|
||||
raise ValueError(
|
||||
f"LoRA pinned weight cache key collision for {cache_key!r}: "
|
||||
f"cached shape={cached_weight.shape}, dtype={cached_weight.dtype}; "
|
||||
f"new shape={weight.shape}, dtype={weight.dtype}."
|
||||
)
|
||||
|
||||
return cached_weight
|
||||
|
||||
def prepare_lora_batch(
|
||||
self,
|
||||
cur_uids: Set[Optional[str]],
|
||||
@@ -694,7 +768,7 @@ class LoRAMemoryPool:
|
||||
assert (
|
||||
buffer_view.shape == weight.shape
|
||||
), f"LoRA buffer shape {buffer_view.shape} does not match weight shape {weight.shape}."
|
||||
buffer_view.copy_(weight, non_blocking=True)
|
||||
copy_weight_into_buffer(buffer_view, weight)
|
||||
|
||||
if uid is None:
|
||||
for i in range(self.num_layer):
|
||||
@@ -748,7 +822,9 @@ class LoRAMemoryPool:
|
||||
logger.warning(msg)
|
||||
|
||||
for layer_id in range(self.num_layer):
|
||||
layer_weights = lora_adapter.layers[layer_id].weights
|
||||
layer = lora_adapter.layers[layer_id]
|
||||
layer_weights = layer.weights
|
||||
pinned_layer_weights = layer.pinned_weights
|
||||
# - Standard: module_name -> torch.Tensor
|
||||
# - MoE: module_name -> Dict[expert_id -> torch.Tensor]
|
||||
temp_A_buffer: Dict[str, Union[torch.Tensor, Dict[int, torch.Tensor]]] = {
|
||||
@@ -757,6 +833,12 @@ class LoRAMemoryPool:
|
||||
temp_B_buffer: Dict[str, Union[torch.Tensor, Dict[int, torch.Tensor]]] = {
|
||||
target_module: None for target_module in self.B_buffer
|
||||
}
|
||||
temp_A_cache_keys: Dict[str, Optional[Union[str, Dict[int, str]]]] = {
|
||||
target_module: None for target_module in self.A_buffer
|
||||
}
|
||||
temp_B_cache_keys: Dict[str, Optional[Union[str, Dict[int, str]]]] = {
|
||||
target_module: None for target_module in self.B_buffer
|
||||
}
|
||||
|
||||
for name, weights in layer_weights.items():
|
||||
target_module = get_target_module_name(name, self.target_modules)
|
||||
@@ -770,25 +852,33 @@ class LoRAMemoryPool:
|
||||
if temp_A_buffer[target_module] is None:
|
||||
temp_A_buffer[target_module] = {}
|
||||
temp_B_buffer[target_module] = {}
|
||||
temp_A_cache_keys[target_module] = {}
|
||||
temp_B_cache_keys[target_module] = {}
|
||||
|
||||
expert_id = int(expert_match.group(1))
|
||||
if "lora_A" in name:
|
||||
temp_A_buffer[target_module][expert_id] = weights
|
||||
temp_A_cache_keys[target_module][expert_id] = name
|
||||
else:
|
||||
temp_B_buffer[target_module][expert_id] = weights
|
||||
temp_B_cache_keys[target_module][expert_id] = name
|
||||
elif "experts" in name and weights.dim() == 3:
|
||||
# Shared outer MoE weight — 3D tensor [expert_dim, rank, hidden]
|
||||
target_module = target_module + "_moe"
|
||||
if "lora_A" in name:
|
||||
temp_A_buffer[target_module] = weights
|
||||
temp_A_cache_keys[target_module] = name
|
||||
else:
|
||||
temp_B_buffer[target_module] = weights
|
||||
temp_B_cache_keys[target_module] = name
|
||||
else:
|
||||
# Standard weight — single tensor per module
|
||||
if "lora_A" in name:
|
||||
temp_A_buffer[target_module] = weights
|
||||
temp_A_cache_keys[target_module] = name
|
||||
else:
|
||||
temp_B_buffer[target_module] = weights
|
||||
temp_B_cache_keys[target_module] = name
|
||||
|
||||
# Track which buffer keys correspond to a real wrapped module on
|
||||
# this layer. `temp_A/B_buffer` is seeded with every key in the
|
||||
@@ -822,6 +912,12 @@ class LoRAMemoryPool:
|
||||
target_module,
|
||||
)
|
||||
)
|
||||
cache_keys = temp_A_cache_keys[target_module]
|
||||
assert cache_keys is not None
|
||||
temp_A_cache_keys[target_module] = append_cache_key_suffix(
|
||||
cache_keys,
|
||||
f"moe_tp{self.moe_tp_rank}",
|
||||
)
|
||||
if temp_B_buffer.get(target_module) is not None:
|
||||
temp_B_buffer[target_module] = (
|
||||
module.slice_moe_lora_b_weights(
|
||||
@@ -830,6 +926,12 @@ class LoRAMemoryPool:
|
||||
target_module,
|
||||
)
|
||||
)
|
||||
cache_keys = temp_B_cache_keys[target_module]
|
||||
assert cache_keys is not None
|
||||
temp_B_cache_keys[target_module] = append_cache_key_suffix(
|
||||
cache_keys,
|
||||
f"moe_tp{self.moe_tp_rank}",
|
||||
)
|
||||
|
||||
continue
|
||||
|
||||
@@ -849,9 +951,22 @@ class LoRAMemoryPool:
|
||||
temp_A_buffer[target_module] = module.slice_lora_a_weights(
|
||||
temp_A_buffer[target_module], self.tp_rank
|
||||
)
|
||||
cache_keys = temp_A_cache_keys[target_module]
|
||||
assert cache_keys is not None
|
||||
temp_A_cache_keys[target_module] = append_cache_key_suffix(
|
||||
cache_keys,
|
||||
f"tp{self.tp_rank}",
|
||||
)
|
||||
|
||||
temp_B_buffer[target_module] = module.slice_lora_b_weights(
|
||||
temp_B_buffer[target_module], self.tp_rank
|
||||
)
|
||||
cache_keys = temp_B_cache_keys[target_module]
|
||||
assert cache_keys is not None
|
||||
temp_B_cache_keys[target_module] = append_cache_key_suffix(
|
||||
cache_keys,
|
||||
f"tp{self.tp_rank}",
|
||||
)
|
||||
|
||||
for name, weights in temp_A_buffer.items():
|
||||
if name not in active_target_modules:
|
||||
@@ -859,6 +974,7 @@ class LoRAMemoryPool:
|
||||
c = get_stacked_multiply(name, self.base_model)
|
||||
max_r = self.max_lora_rank
|
||||
target_buffer = self.A_buffer[name][layer_id]
|
||||
weights_cache_key = temp_A_cache_keys[name]
|
||||
|
||||
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
||||
if self.experts_shared_outer_loras and name == "gate_up_proj_moe":
|
||||
@@ -875,6 +991,12 @@ class LoRAMemoryPool:
|
||||
f"gate_up_proj_moe lora_A has expert_dim="
|
||||
f"{weights.shape[0]} (expected 1)."
|
||||
)
|
||||
assert isinstance(weights_cache_key, str)
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
weights_cache_key,
|
||||
weights,
|
||||
)
|
||||
representative_weight = weights[0]
|
||||
buffer_view = target_buffer[
|
||||
buffer_id, 0, : lora_rank * c, :
|
||||
@@ -888,6 +1010,11 @@ class LoRAMemoryPool:
|
||||
f"{len(weights)} entries (expected 1)."
|
||||
)
|
||||
rep = next(iter(weights.values()))
|
||||
assert isinstance(weights_cache_key, dict)
|
||||
rep_cache_key = next(iter(weights_cache_key.values()))
|
||||
rep = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights, rep_cache_key, rep
|
||||
)
|
||||
representative_weight = rep
|
||||
buffer_view = target_buffer[
|
||||
buffer_id, 0, : lora_rank * c, :
|
||||
@@ -921,11 +1048,21 @@ class LoRAMemoryPool:
|
||||
# the MoE kernel's [:max_r] / [max_r:2*max_r] slicing
|
||||
# is correct.
|
||||
target_buffer[buffer_id].zero_()
|
||||
for local_eid, expert_weight in self._iter_local_expert_weights(
|
||||
weights
|
||||
assert isinstance(weights_cache_key, (str, dict))
|
||||
for (
|
||||
local_eid,
|
||||
expert_weight,
|
||||
expert_cache_key,
|
||||
) in self._iter_local_expert_weights(
|
||||
weights, weights_cache_key
|
||||
):
|
||||
if expert_weight is None:
|
||||
continue
|
||||
expert_weight = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
expert_cache_key,
|
||||
expert_weight,
|
||||
)
|
||||
for ci in range(c):
|
||||
buffer_view = target_buffer[
|
||||
buffer_id,
|
||||
@@ -941,12 +1078,20 @@ class LoRAMemoryPool:
|
||||
)
|
||||
else:
|
||||
buffer_view = target_buffer[buffer_id, : lora_rank * c, :]
|
||||
if weights is not None:
|
||||
assert isinstance(weights_cache_key, str)
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
weights_cache_key,
|
||||
weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, weights)
|
||||
|
||||
for name, weights in temp_B_buffer.items():
|
||||
if name not in active_target_modules:
|
||||
continue
|
||||
target_buffer = self.B_buffer[name][layer_id]
|
||||
weights_cache_key = temp_B_cache_keys[name]
|
||||
|
||||
if name in ["gate_up_proj_moe", "down_proj_moe"]:
|
||||
if self.experts_shared_outer_loras and name == "down_proj_moe":
|
||||
@@ -962,11 +1107,17 @@ class LoRAMemoryPool:
|
||||
)
|
||||
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
|
||||
w = weights[0]
|
||||
assert isinstance(weights_cache_key, str)
|
||||
if w is not None:
|
||||
w = w * lora_adapter.scaling
|
||||
w = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
append_cache_key_suffix(
|
||||
weights_cache_key, "expert0"
|
||||
),
|
||||
w,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, w)
|
||||
# Zero beyond loaded rank — MoE kernel reads full max_rank
|
||||
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
|
||||
elif isinstance(weights, dict) and len(weights) > 0:
|
||||
if len(weights) != 1:
|
||||
raise ValueError(
|
||||
@@ -975,37 +1126,65 @@ class LoRAMemoryPool:
|
||||
f"{len(weights)} entries (expected 1)."
|
||||
)
|
||||
rep = next(iter(weights.values()))
|
||||
assert isinstance(weights_cache_key, dict)
|
||||
rep_cache_key = next(iter(weights_cache_key.values()))
|
||||
buffer_view = target_buffer[buffer_id, 0, :, :lora_rank]
|
||||
if rep is not None:
|
||||
rep = rep * lora_adapter.scaling
|
||||
rep = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
rep_cache_key,
|
||||
rep,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, rep)
|
||||
# Zero beyond loaded rank — MoE kernel reads full max_rank
|
||||
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unexpected weight format for shared outer down_proj_moe lora_B: "
|
||||
f"type={type(weights)}, "
|
||||
f"shape={weights.shape if isinstance(weights, torch.Tensor) else 'N/A'}"
|
||||
)
|
||||
# Zero beyond loaded rank — MoE kernel reads full max_rank.
|
||||
target_buffer[buffer_id, 0, :, lora_rank:].zero_()
|
||||
elif isinstance(weights, (torch.Tensor, dict)):
|
||||
# Zero out slots this rank owns but the adapter
|
||||
# doesn't fill (padded-out / out-of-rank experts);
|
||||
# then scale+load the ones it does.
|
||||
target_buffer[buffer_id].zero_()
|
||||
for local_eid, w in self._iter_local_expert_weights(weights):
|
||||
assert isinstance(weights_cache_key, (str, dict))
|
||||
for (
|
||||
local_eid,
|
||||
w,
|
||||
w_cache_key,
|
||||
) in self._iter_local_expert_weights(
|
||||
weights, weights_cache_key
|
||||
):
|
||||
if w is not None:
|
||||
w = w * lora_adapter.scaling
|
||||
w = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
w_cache_key,
|
||||
w,
|
||||
)
|
||||
buffer_view = target_buffer[
|
||||
buffer_id, local_eid, :, :lora_rank
|
||||
]
|
||||
load_lora_weight_tensor(buffer_view, w)
|
||||
else:
|
||||
buffer_view = target_buffer[buffer_id, :, :lora_rank]
|
||||
if weights is not None:
|
||||
assert isinstance(weights_cache_key, str)
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_layer_weights,
|
||||
weights_cache_key,
|
||||
weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, weights)
|
||||
|
||||
if lora_adapter.embedding_layers:
|
||||
org_vocab_size = self.base_hf_config.vocab_size
|
||||
lora_added_tokens_size = lora_adapter.config.lora_added_tokens_size
|
||||
pinned_embedding_layers = lora_adapter.pinned_embedding_layers
|
||||
pinned_added_tokens_embeddings = lora_adapter.pinned_added_tokens_embeddings
|
||||
# Only when LoRA is applied to the embedding layer will it have the extra-token issue that needs to be resolved.
|
||||
# Load embeddings weights for extra tokens to buffer
|
||||
if lora_adapter.added_tokens_embeddings:
|
||||
@@ -1014,6 +1193,11 @@ class LoRAMemoryPool:
|
||||
buffer_view = self.new_embeddings_buffer["input_embeddings"][
|
||||
buffer_id, :lora_added_tokens_size
|
||||
]
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_added_tokens_embeddings,
|
||||
name,
|
||||
weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, weights)
|
||||
|
||||
# load vocab_emb and lm_head
|
||||
@@ -1029,6 +1213,11 @@ class LoRAMemoryPool:
|
||||
:lora_rank,
|
||||
: (org_vocab_size + lora_added_tokens_size),
|
||||
]
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_embedding_layers,
|
||||
name,
|
||||
weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, weights)
|
||||
elif (
|
||||
target_module == "embed_tokens"
|
||||
@@ -1042,6 +1231,11 @@ class LoRAMemoryPool:
|
||||
buffer_view = self.embedding_B_buffer[target_module][
|
||||
buffer_id, :, :lora_rank
|
||||
]
|
||||
lora_b_weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_embedding_layers,
|
||||
name,
|
||||
lora_b_weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, lora_b_weights)
|
||||
|
||||
elif (
|
||||
@@ -1056,6 +1250,11 @@ class LoRAMemoryPool:
|
||||
:lora_rank,
|
||||
:,
|
||||
]
|
||||
weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_embedding_layers,
|
||||
name,
|
||||
weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, weights)
|
||||
elif (
|
||||
target_module == "lm_head"
|
||||
@@ -1070,12 +1269,20 @@ class LoRAMemoryPool:
|
||||
lora_b_weights = lora_lm_head_module.slice_lora_b_weights(
|
||||
lora_b_weights, self.tp_rank
|
||||
)
|
||||
cache_key = append_cache_key_suffix(name, f"tp{self.tp_rank}")
|
||||
else:
|
||||
cache_key = name
|
||||
|
||||
buffer_view = self.lm_head_B_buffer[target_module][
|
||||
buffer_id,
|
||||
: lora_b_weights.shape[0],
|
||||
:lora_rank,
|
||||
]
|
||||
lora_b_weights = self._get_maybe_cached_weight_for_transfer(
|
||||
pinned_embedding_layers,
|
||||
cache_key,
|
||||
lora_b_weights,
|
||||
)
|
||||
load_lora_weight_tensor(buffer_view, lora_b_weights)
|
||||
elif (
|
||||
target_module == "lm_head"
|
||||
|
||||
@@ -83,6 +83,26 @@ class LoRAType(Enum):
|
||||
LORA_B = 1
|
||||
|
||||
|
||||
def copy_weight_into_buffer(
|
||||
buffer_view: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
) -> None:
|
||||
"""
|
||||
Copy a LoRA weight tensor into a destination buffer.
|
||||
|
||||
When a pinned CPU source has a dtype mismatch with a device destination,
|
||||
cast on the destination device instead of doing the conversion on CPU.
|
||||
"""
|
||||
if weight.dtype == buffer_view.dtype:
|
||||
buffer_view.copy_(weight, non_blocking=True)
|
||||
return
|
||||
|
||||
if weight.device.type == "cpu" and buffer_view.device.type != "cpu":
|
||||
weight = weight.to(device=buffer_view.device, non_blocking=True)
|
||||
|
||||
buffer_view.copy_(weight.to(dtype=buffer_view.dtype), non_blocking=True)
|
||||
|
||||
|
||||
def get_hidden_dim(
|
||||
module_name: str,
|
||||
config: AutoConfig,
|
||||
|
||||
Reference in New Issue
Block a user