diff --git a/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b.sh b/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b.sh new file mode 100755 index 0000000000..8da7296e0d --- /dev/null +++ b/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b.sh @@ -0,0 +1,68 @@ +#!/usr/bin/env bash + +set -euo pipefail +set -x + +# Colocated GRPO for Qwen3-0.6B on GSM8K with Megatron and LoRA-XS. + +LORA_RANK="${LORA_RANK:-32}" +LORA_XS_INIT_DIR="${LORA_XS_INIT_DIR:-$HOME/lora_xs/qwen3_0.6b_lora_xs_r${LORA_RANK}}" +SOURCE_MODEL="${SOURCE_MODEL:-Qwen/Qwen3-0.6B}" +POLICY_MODEL="${POLICY_MODEL:-$SOURCE_MODEL}" +DATA_DIR="${DATA_DIR:-$HOME/data/gsm8k}" +NUM_GPUS="${NUM_GPUS:-8}" +LOGGER="${LOGGER:-wandb}" +LORA_ALPHA="${LORA_ALPHA:-$LORA_RANK}" + +uv run --isolated --extra megatron -m skyrl.train.entrypoints.main_base \ + data.train_data="['$DATA_DIR/train.parquet']" \ + data.val_data="['$DATA_DIR/validation.parquet']" \ + trainer.algorithm.advantage_estimator=grpo \ + trainer.policy.model.path="$POLICY_MODEL" \ + trainer.ref.model.path="$SOURCE_MODEL" \ + trainer.placement.colocate_all=true \ + trainer.strategy=megatron \ + trainer.placement.policy_num_gpus_per_node="$NUM_GPUS" \ + trainer.placement.ref_num_gpus_per_node="$NUM_GPUS" \ + generator.inference_engine.num_engines="$NUM_GPUS" \ + generator.inference_engine.tensor_parallel_size=1 \ + trainer.policy.megatron_config.tensor_model_parallel_size=1 \ + trainer.policy.megatron_config.pipeline_model_parallel_size=1 \ + trainer.policy.megatron_config.context_parallel_size=1 \ + trainer.policy.megatron_config.lora_config.lora_type=lora_xs \ + trainer.ref.megatron_config.tensor_model_parallel_size=1 \ + trainer.ref.megatron_config.pipeline_model_parallel_size=1 \ + trainer.ref.megatron_config.context_parallel_size=1 \ + trainer.policy.model.lora.rank="$LORA_RANK" \ + trainer.policy.model.lora.alpha="$LORA_ALPHA" \ + trainer.policy.model.lora.target_modules=all-linear \ + trainer.gradient_checkpointing=true \ + trainer.remove_microbatch_padding=true \ + trainer.epochs=20 \ + trainer.eval_batch_size=1024 \ + trainer.eval_before_train=false \ + trainer.eval_interval=5 \ + trainer.update_epochs_per_batch=1 \ + trainer.train_batch_size=128 \ + trainer.policy_mini_batch_size=64 \ + trainer.micro_forward_batch_size_per_gpu=4 \ + trainer.micro_train_batch_size_per_gpu=4 \ + trainer.ckpt_interval=10 \ + trainer.max_prompt_length=512 \ + generator.sampling_params.max_generate_length=1024 \ + trainer.policy.optimizer_config.lr=1.0e-5 \ + trainer.algorithm.use_kl_loss=true \ + generator.inference_engine.backend=vllm \ + generator.inference_engine.run_engines_locally=true \ + generator.inference_engine.weight_sync_backend=nccl \ + generator.batched=true \ + generator.n_samples_per_prompt=5 \ + generator.inference_engine.gpu_memory_utilization=0.6 \ + environment.env_class=gsm8k \ + trainer.logger="$LOGGER" \ + trainer.project_name=gsm8k_megatron \ + trainer.run_name="gsm8k_megatron_qwen3_0.6b_lora_xs_r${LORA_RANK}_a${LORA_ALPHA}" \ + trainer.resume_mode=from_path \ + trainer.resume_path="$LORA_XS_INIT_DIR/global_step_0" \ + trainer.ckpt_path="$HOME/ckpts/gsm8k_0.6b_lora_xs_ckpt" \ + "$@" diff --git a/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b_producer.sh b/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b_producer.sh new file mode 100755 index 0000000000..cc7beb0192 --- /dev/null +++ b/examples/train/lora_xs/run_megatron_lora_xs_qwen3-0.6b_producer.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash + +set -euo pipefail + +BASE_MODEL="${BASE_MODEL:-Qwen/Qwen3-0.6B}" +LORA_RANK="${LORA_RANK:-32}" +LORA_XS_INIT_DIR="${LORA_XS_INIT_DIR:-$HOME/lora_xs/qwen3_0.6b_lora_xs_r${LORA_RANK}}" +DEFAULT_CONFIG_OVERRIDES='{"trainer.placement.policy_num_gpus_per_node":8,"trainer.policy.megatron_config.tensor_model_parallel_size":1,"trainer.policy.megatron_config.pipeline_model_parallel_size":1,"trainer.policy.megatron_config.context_parallel_size":1}' +CONFIG_OVERRIDES="${CONFIG_OVERRIDES:-$DEFAULT_CONFIG_OVERRIDES}" + +uv run --isolated --extra megatron -m skyrl.train.entrypoints.lora_xs_init \ + --base-model "$BASE_MODEL" \ + --rank "$LORA_RANK" \ + --output-dir "$LORA_XS_INIT_DIR" \ + --config-overrides "$CONFIG_OVERRIDES" diff --git a/skyrl/backends/skyrl_train/distributed/megatron/megatron_strategy.py b/skyrl/backends/skyrl_train/distributed/megatron/megatron_strategy.py index a4dabcb7dc..98f6046266 100644 --- a/skyrl/backends/skyrl_train/distributed/megatron/megatron_strategy.py +++ b/skyrl/backends/skyrl_train/distributed/megatron/megatron_strategy.py @@ -105,6 +105,7 @@ def _patched_update_fp32_params_by_new_state(self): _orig_load_parameter_state_from_dp_reshardable = DistributedOptimizer.load_parameter_state_from_dp_reshardable +@torch.no_grad() def _patched_load_parameter_state_from_dp_reshardable(self, state_dict): """Wrapper around the original method that preserves the Adam step counter. diff --git a/skyrl/backends/skyrl_train/distributed/megatron/optimizer.py b/skyrl/backends/skyrl_train/distributed/megatron/optimizer.py index 5c2a99318d..910a6d8eae 100644 --- a/skyrl/backends/skyrl_train/distributed/megatron/optimizer.py +++ b/skyrl/backends/skyrl_train/distributed/megatron/optimizer.py @@ -33,7 +33,7 @@ def init_megatron_optim_config( - optim_config: Union[SkyRLOptimizerConfig, DictConfig], optimizer_config_kwargs: dict + optim_config: Union[SkyRLOptimizerConfig, DictConfig], optimizer_config_kwargs: dict, bf16: bool = True ) -> OptimizerConfig: adam_betas = getattr(optim_config, "adam_betas", (0.9, 0.999)) optim_args = { @@ -44,8 +44,9 @@ def init_megatron_optim_config( "weight_decay": getattr(optim_config, "weight_decay", 1e-2), "adam_beta1": float(adam_betas[0]), "adam_beta2": float(adam_betas[1]), - "bf16": True, - "params_dtype": torch.bfloat16, + "bf16": bf16, + "params_dtype": torch.bfloat16 if bf16 else torch.float32, + "store_param_remainders": bf16, "use_distributed_optimizer": True, } # YAML dtype overrides arrive as strings; Megatron expects torch.dtype. diff --git a/skyrl/backends/skyrl_train/workers/megatron/lora_xs.py b/skyrl/backends/skyrl_train/workers/megatron/lora_xs.py new file mode 100644 index 0000000000..24aebe5c4d --- /dev/null +++ b/skyrl/backends/skyrl_train/workers/megatron/lora_xs.py @@ -0,0 +1,90 @@ +"""LoRA-XS adapters for Megatron training and PEFT export.""" + +import contextlib + +import torch +from megatron.bridge.peft.lora import LoRA +from megatron.bridge.peft.lora_layers import LoRALinear, TEFusedLoRALinear +from megatron.bridge.peft.utils import ParallelLinearAdapter +from torch import nn + +LORA_XS_INIT_STD = 1e-5 + + +class LoRAXSCore(nn.Linear): + """Trainable square projection between frozen LoRA factors.""" + + +def configure_lora_xs_adapter(adapter: ParallelLinearAdapter) -> None: + """Freeze LoRA factors and insert a noise-initialized square core.""" + if adapter.is_expert: + raise ValueError("LoRA-XS does not support expert adapters") + if not isinstance(adapter.activation, nn.Identity): + raise TypeError("LoRA-XS requires identity LoRA activation") + + adapter.linear_in.weight.requires_grad_(False) + adapter.linear_out.weight.requires_grad_(False) + core = LoRAXSCore( + adapter.dim, + adapter.dim, + bias=False, + device=adapter.linear_in.weight.device, + dtype=adapter.linear_in.weight.dtype, + ) + nn.init.normal_(core.weight, std=LORA_XS_INIT_STD) + adapter.activation = core + + +class LoRAXS(LoRA): + """LoRA with frozen SVD factors and a trainable square core.""" + + def export_context(self, model_chunks): + """Materialize standard LoRA weights while exporting.""" + return materialized_lora_xs_adapters(model_chunks) + + def transform(self, module: nn.Module, name=None, prefix=None) -> nn.Module: + if isinstance(module, LoRALinear): + return module + + transformed = super().transform(module, name, prefix) + if transformed is module: + return module + if isinstance(transformed, TEFusedLoRALinear): + transformed = LoRALinear(transformed.to_wrap, transformed.adapter) + if not isinstance(transformed, LoRALinear) or not isinstance(transformed.adapter, ParallelLinearAdapter): + raise TypeError("LoRA-XS supports only Megatron parallel linear adapters") + + configure_lora_xs_adapter(transformed.adapter) + return transformed + + +def _get_lora_xs_adapters(model_chunks) -> list[ParallelLinearAdapter]: + chunks = model_chunks if isinstance(model_chunks, (list, tuple)) else [model_chunks] + return [ + module.adapter + for chunk in chunks + for module in chunk.modules() + if isinstance(module, LoRALinear) + and isinstance(module.adapter, ParallelLinearAdapter) + and isinstance(module.adapter.activation, LoRAXSCore) + ] + + +@contextlib.contextmanager +def materialized_lora_xs_adapters(model_chunks): + """Temporarily expose LoRA-XS adapters as standard two-factor LoRA.""" + saved = [] + try: + with torch.no_grad(): + for adapter in _get_lora_xs_adapters(model_chunks): + core = adapter.activation + linear_out = adapter.linear_out.weight + saved.append((adapter, core, linear_out.detach().clone())) + linear_out.copy_(linear_out.float() @ core.weight.float()) + adapter.activation = nn.Identity() + yield + finally: + with torch.no_grad(): + for adapter, core, linear_out in saved: + adapter.linear_out.weight.copy_(linear_out) + adapter.activation = core diff --git a/skyrl/backends/skyrl_train/workers/megatron/lora_xs_init_worker.py b/skyrl/backends/skyrl_train/workers/megatron/lora_xs_init_worker.py new file mode 100644 index 0000000000..88386bc6c7 --- /dev/null +++ b/skyrl/backends/skyrl_train/workers/megatron/lora_xs_init_worker.py @@ -0,0 +1,144 @@ +"""Megatron worker for offline LoRA-XS initialization.""" + +import ray +import torch +from loguru import logger +from megatron.bridge.peft.lora_layers import LoRALinear +from megatron.bridge.peft.utils import ParallelLinearAdapter +from megatron.core.distributed.distributed_data_parallel_config import ( + DistributedDataParallelConfig, +) + +from skyrl.backends.skyrl_train.workers.megatron.lora_xs import ( + LORA_XS_INIT_STD, + LoRAXSCore, +) +from skyrl.backends.skyrl_train.workers.megatron.megatron_worker import ( + MegatronPolicyWorkerBase, +) +from skyrl.backends.skyrl_train.workers.megatron.pissa_init_worker import ( + _all_gather, + _shard, +) +from skyrl.train.config.config import get_config_as_dict + + +def lora_xs_factors(weight: torch.Tensor, rank: int) -> tuple[torch.Tensor, torch.Tensor]: + """Return the principal right vectors and singular-value-scaled left vectors.""" + max_rank = min(weight.shape) + if rank > max_rank: + raise ValueError(f"LoRA-XS rank {rank} exceeds maximum rank {max_rank} for weight shape {tuple(weight.shape)}") + + u, s, vh = torch.linalg.svd(weight.float(), full_matrices=False) + linear_in = (s[:rank].unsqueeze(1) * vh[:rank, :]).contiguous() + linear_out = u[:, :rank].contiguous() + return linear_in, linear_out + + +def _synchronized_lora_xs_factors(weight: torch.Tensor, rank: int, group) -> tuple[torch.Tensor, torch.Tensor]: + if torch.distributed.get_rank(group) == 0: + linear_in, linear_out = lora_xs_factors(weight, rank) + else: + linear_in = torch.empty((rank, weight.shape[1]), dtype=torch.float32, device=weight.device) + linear_out = torch.empty((weight.shape[0], rank), dtype=torch.float32, device=weight.device) + + torch.distributed.broadcast(linear_in, group=group, group_src=0) + torch.distributed.broadcast(linear_out, group=group, group_src=0) + return linear_in, linear_out + + +@torch.no_grad() +def _init_one_lora_xs_adapter( + base_linear, + adapter, + tp_size: int, + tp_rank: int, + tp_group, +) -> None: + base_weight = base_linear.weight + if base_weight.is_meta: + raise RuntimeError("LoRA-XS requires pretrained weights before adapter initialization") + + base_shard_dim = 1 if adapter.input_is_parallel else 0 + if tp_size == 1: + full_weight = base_weight.detach().float() + linear_in, linear_out = lora_xs_factors(full_weight, adapter.dim) + else: + full_weight = _all_gather(base_weight.detach().float(), base_shard_dim, tp_size, tp_group) + linear_in, linear_out = _synchronized_lora_xs_factors(full_weight, adapter.dim, tp_group) + + linear_in_shard_dim = 1 if adapter.input_is_parallel else 0 + dtype = adapter.linear_in.weight.dtype + adapter.linear_in.weight.copy_(_shard(linear_in, linear_in_shard_dim, tp_rank, tp_size).to(dtype)) + adapter.linear_out.weight.copy_(_shard(linear_out, 0, tp_rank, tp_size).to(dtype)) + adapter.activation.weight.normal_(std=LORA_XS_INIT_STD) + + +def apply_lora_xs_init(model_chunks) -> None: + """Initialize dense LoRA-XS adapters from their base weights.""" + import megatron.core.parallel_state as mpu + + tp_size = mpu.get_tensor_model_parallel_world_size() + tp_rank = mpu.get_tensor_model_parallel_rank() + tp_group = mpu.get_tensor_model_parallel_group() + chunks = model_chunks if isinstance(model_chunks, (list, tuple)) else [model_chunks] + adapters = [] + for chunk in chunks: + for module in chunk.modules(): + if isinstance(module, LoRALinear) and isinstance(module.adapter, ParallelLinearAdapter): + if not isinstance(module.adapter.activation, LoRAXSCore): + raise TypeError("LoRA-XS found a non-LoRA-XS parallel adapter") + adapters.append((module.to_wrap, module.adapter)) + + if not adapters: + raise ValueError("LoRA-XS found no supported adapters") + for base_linear, adapter in adapters: + _init_one_lora_xs_adapter( + base_linear, + adapter, + tp_size, + tp_rank, + tp_group, + ) + logger.info(f"LoRA-XS: initialized {len(adapters)} adapter(s)") + + +class LoRAXSInitWorkerBase(MegatronPolicyWorkerBase): + def make_megatron_module( + self, + wrap_with_ddp=True, + ddp_config=None, + lora_config=None, + lora_type="lora_xs", + bf16=True, + ): + if lora_type != "lora_xs": + raise ValueError("LoRA-XS initialization requires lora_type='lora_xs'") + self.configure_lora(lora_config, lora_type) + + def lora_pre_wrap_hook(model): + lora_model = self.lora_cls(model, training=True) + self.lora_cls.set_params_to_save(lora_model) + return lora_model + + def lora_xs_pre_wrap_hook(model): + apply_lora_xs_init(model) + return model + + self.provider.register_pre_wrap_hook(lora_pre_wrap_hook) + self.provider.register_pre_wrap_hook(lora_xs_pre_wrap_hook) + + resolved_ddp_config = DistributedDataParallelConfig() + if wrap_with_ddp: + resolved_ddp_config.use_distributed_optimizer = True + if ddp_config is not None: + for key, value in get_config_as_dict(ddp_config).items(): + setattr(resolved_ddp_config, key, value) + return self.provider.provide_distributed_model( + ddp_config=resolved_ddp_config, + wrap_with_ddp=wrap_with_ddp, + bf16=bf16, + ) + + +LoRAXSInitWorker = ray.remote(num_gpus=1)(LoRAXSInitWorkerBase) diff --git a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py index 7aa0511a7a..76ce81eb7c 100644 --- a/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py +++ b/skyrl/backends/skyrl_train/workers/megatron/megatron_worker.py @@ -1,6 +1,7 @@ import os import shutil from collections import defaultdict +from contextlib import nullcontext from datetime import timedelta from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union @@ -58,6 +59,7 @@ LoraSignature, iter_opts, ) +from skyrl.backends.skyrl_train.workers.megatron.lora_xs import LoRAXS from skyrl.backends.skyrl_train.workers.megatron.megatron_model_wrapper import ( MegatronModelWrapper, ) @@ -467,7 +469,10 @@ def init_configs( provider.tensor_model_parallel_size = megatron_config.tensor_model_parallel_size provider.pipeline_model_parallel_size = megatron_config.pipeline_model_parallel_size - provider.pipeline_dtype = torch.bfloat16 if bf16 else torch.float32 + provider.params_dtype = torch.bfloat16 if bf16 else torch.float32 + provider.bf16 = bf16 + provider.fp16 = False + provider.pipeline_dtype = provider.params_dtype provider.context_parallel_size = megatron_config.context_parallel_size provider.expert_model_parallel_size = megatron_config.expert_model_parallel_size provider.expert_tensor_parallel_size = megatron_config.expert_tensor_parallel_size @@ -583,6 +588,21 @@ def configure_lora(self, lora_config, lora_type: Optional[str] = "lora"): lora_B_init_method="zero", exclude_modules=[] if lora_config.exclude_modules is None else lora_config.exclude_modules, ) + elif lora_type == "lora_xs": + self.lora_cls = LoRAXS( + target_modules=( + ["linear_qkv", "linear_proj", "linear_fc1", "linear_fc2"] + if lora_config.target_modules == "all-linear" + else lora_config.target_modules + ), + dim=lora_config.rank, + alpha=lora_config.alpha, + dropout=lora_config.dropout, + lora_A_init_method=lora_config.init_method, + lora_B_init_method="zero", + exclude_modules=[] if lora_config.exclude_modules is None else lora_config.exclude_modules, + lora_dtype=torch.bfloat16 if self.cfg.bf16 else torch.float32, + ) def make_megatron_module( self, @@ -962,7 +982,9 @@ def init_model(self, model_path, num_training_steps: int = 1e9): self.scheduler = None else: optim_config = init_megatron_optim_config( - self.cfg.policy.optimizer_config, self.cfg.policy.megatron_config.optimizer_config_kwargs + self.cfg.policy.optimizer_config, + self.cfg.policy.megatron_config.optimizer_config_kwargs, + bf16=self.cfg.bf16, ) self.optimizer = get_megatron_optimizer(self.actor_module, optim_config) @@ -1416,8 +1438,10 @@ async def _save_lora_adapters_and_sync( from safetensors.torch import save_file adapter_state = {} - for name, tensor in self.bridge.export_adapter_weights(self.actor_module, cpu=True, show_progress=False): - adapter_state[f"base_model.model.{name}"] = tensor.clone().float() + export_context = getattr(self.lora_cls, "export_context", nullcontext) + with export_context(self.actor_module): + for name, tensor in self.bridge.export_adapter_weights(self.actor_module, cpu=True, show_progress=False): + adapter_state[f"base_model.model.{name}"] = tensor.clone().float() if torch.distributed.get_rank() == 0: os.makedirs(lora_sync_path, exist_ok=True) diff --git a/skyrl/backends/skyrl_train_backend.py b/skyrl/backends/skyrl_train_backend.py index ea08b13d01..cd63256130 100644 --- a/skyrl/backends/skyrl_train_backend.py +++ b/skyrl/backends/skyrl_train_backend.py @@ -20,7 +20,10 @@ VLLMRenderer, render_model_input, ) -from skyrl.backends.skyrl_train.inference_servers.utils import resolve_policy_model_name +from skyrl.backends.skyrl_train.inference_servers.utils import ( + _uses_lora_weight_sync, + resolve_policy_model_name, +) from skyrl.backends.skyrl_train.training_batch import ( TensorList, TrainingInputBatch, @@ -1045,11 +1048,11 @@ def _sample_with_remote_client( # Resolve the inference-engine model name per request. With multi-LoRA # the adapter name on vLLM IS the Tinker model_id (registered by - # save_sampler_checkpoint via load_lora_adapter). Single-tenant / - # FFT path falls back to resolve_policy_model_name(cfg). + # save_sampler_checkpoint via load_lora_adapter). Full-weight sync and + # FFT fall back to resolve_policy_model_name(cfg). fallback_model_name = resolve_policy_model_name(self._cfg) per_request_models = [ - mid if (self._base_lora_signature is not None and mid in self._model_ids_to_role) else fallback_model_name + mid if (_uses_lora_weight_sync(self._cfg) and mid in self._model_ids_to_role) else fallback_model_name for mid in prepared_batch.all_model_ids ] diff --git a/skyrl/tinker/extra/skyrl_train_inference_forwarding.py b/skyrl/tinker/extra/skyrl_train_inference_forwarding.py index 227db2de92..27141bd078 100644 --- a/skyrl/tinker/extra/skyrl_train_inference_forwarding.py +++ b/skyrl/tinker/extra/skyrl_train_inference_forwarding.py @@ -17,6 +17,25 @@ from skyrl.tinker.db_models import EngineStateDB, FutureDB, RequestStatus from skyrl.utils.log import logger +_MEGATRON_MERGE_LORA_KEY = "trainer.policy.megatron_config.lora_config.merge_lora" +_SERVED_MODEL_NAME_KEY = "generator.inference_engine.served_model_name" + + +def _resolve_forwarded_model_name( + engine_config: EngineConfig, + model_id: str, + base_model: str | None, +) -> str: + """Resolve the vLLM model name for an API-forwarded sample.""" + if base_model: + return base_model + + backend_config = engine_config.backend_config + merge_lora = engine_config.backend == "megatron" and backend_config.get(_MEGATRON_MERGE_LORA_KEY, True) + if merge_lora: + return backend_config.get(_SERVED_MODEL_NAME_KEY) or engine_config.base_model + return model_id + class SkyRLTrainInferenceForwardingClient: """Forwards EXTERNAL sample requests to the SkyRL-Train-managed vLLM.""" @@ -112,9 +131,7 @@ async def _forward_with_retry(self, sample_req, model_id: str, *, base_model: st async def _forward( self, proxy_url: str, sample_req, model_id: str, *, base_model: str | None ) -> types.SampleOutput: - # model_id matches the LoRA name registered with vLLM during - # save_weights_for_sampler; base_model is used for non-LoRA sampling. - model_name = base_model if base_model else model_id + model_name = _resolve_forwarded_model_name(self.engine_config, model_id, base_model) model_input = sample_req.prompt.to_types() prompt_tokens = render_model_input([model_input])[0].prompt_ids diff --git a/skyrl/train/config/config.py b/skyrl/train/config/config.py index 042153c1b6..1cfa0a9708 100644 --- a/skyrl/train/config/config.py +++ b/skyrl/train/config/config.py @@ -412,7 +412,7 @@ def __post_init__(self) -> None: @dataclass class MegatronLoraConfig(BaseConfig): lora_type: str = "lora" - """``"lora"`` or ``"canonical_lora"``. + """``"lora"``, ``"canonical_lora"``, or ``"lora_xs"``. See https://docs.nvidia.com/nemo/megatron-bridge/0.2.0/apidocs/bridge/bridge.peft.lora.html""" merge_lora: bool = True """Merge LoRA weights into the base weights during weight sync.""" diff --git a/skyrl/train/entrypoints/lora_xs_init.py b/skyrl/train/entrypoints/lora_xs_init.py new file mode 100644 index 0000000000..b204c48f8f --- /dev/null +++ b/skyrl/train/entrypoints/lora_xs_init.py @@ -0,0 +1,88 @@ +"""Produce a step-zero LoRA-XS adapter checkpoint with Megatron.""" + +import argparse +import json +import sys +from pathlib import Path + +import ray +import torch + +from skyrl.backends.skyrl_train.workers.megatron.lora_xs_init_worker import ( + LoRAXSInitWorker, +) +from skyrl.backends.skyrl_train.workers.worker import PPORayActorGroup +from skyrl.train.config import SkyRLTrainConfig, get_config_as_dict +from skyrl.train.utils.utils import initialize_ray +from skyrl.utils.log import logger +from skyrl.utils.tok import get_tokenizer + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base-model", required=True) + parser.add_argument("--rank", required=True, type=int) + parser.add_argument("--output-dir", required=True, type=Path) + parser.add_argument("--config-overrides", type=json.loads, default={}) + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + if args.rank <= 0: + raise ValueError("--rank must be positive") + args.output_dir.mkdir(parents=True, exist_ok=False) + + overrides = dict(args.config_overrides) + overrides.update( + { + "trainer.strategy": "megatron", + "trainer.policy.model.path": args.base_model, + "trainer.policy.model.lora.rank": args.rank, + "trainer.policy.model.lora.alpha": args.rank, + "trainer.policy.model.lora.target_modules": "all-linear", + "trainer.policy.model.lora.exclude_modules": None, + "trainer.policy.megatron_config.lora_config.lora_type": "lora_xs", + "trainer.placement.colocate_all": False, + } + ) + cfg = SkyRLTrainConfig.from_cli_overrides(overrides) + tokenizer = get_tokenizer(args.base_model) + + try: + initialize_ray(cfg) + policy = PPORayActorGroup( + cfg.trainer, + cfg.trainer.placement.policy_num_nodes, + cfg.trainer.placement.policy_num_gpus_per_node, + LoRAXSInitWorker, + num_gpus_per_actor=1, + colocate_all=False, + sequence_parallel_size=cfg.trainer.policy.sequence_parallel_size, + record_memory=cfg.trainer.policy.record_memory, + ) + ray.get(policy.async_init_model(args.base_model, num_training_steps=sys.maxsize)) + ray.get(policy.async_run_ray_method("pass_through", "_set_pad_token_id", tokenizer.pad_token_id)) + ray.get(policy.async_run_ray_method("pass_through", "prime_optimizer_state")) + + checkpoint_dir = args.output_dir / "global_step_0" + ray.get( + policy.async_run_ray_method( + "pass_through", + "save_checkpoint", + str(checkpoint_dir / "policy"), + tokenizer, + ) + ) + torch.save( + {"global_step": 0, "config": get_config_as_dict(cfg)}, + checkpoint_dir / "trainer_state.pt", + ) + logger.info(f"LoRA-XS initialization artifacts saved to {args.output_dir}") + finally: + if ray.is_initialized(): + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_lora_xs.py b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_lora_xs.py new file mode 100644 index 0000000000..ab2d46b3d3 --- /dev/null +++ b/tests/backends/skyrl_train/gpu/gpu_ci/megatron/test_lora_xs.py @@ -0,0 +1,155 @@ +"""Tests for LoRA-XS initialization and standard-LoRA export.""" + +from types import SimpleNamespace + +import pytest +import torch +from megatron.bridge.peft.lora_layers import LoRALinear +from torch import nn + +from skyrl.backends.skyrl_train.workers.megatron import lora_xs, lora_xs_init_worker + +pytestmark = pytest.mark.megatron + + +class _Adapter(nn.Module): + def __init__(self, in_features=8, out_features=12, rank=4): + super().__init__() + self.dim = rank + self.alpha = rank + self.is_expert = False + self.input_is_parallel = False + self.linear_in = nn.Linear(in_features, rank, bias=False) + self.linear_out = nn.Linear(rank, out_features, bias=False) + self.activation = nn.Identity() + + def forward(self, x): + return self.linear_out(self.activation(self.linear_in(x))) + + +def _principal(weight, rank): + u, s, vh = torch.linalg.svd(weight.float(), full_matrices=False) + return (u[:, :rank] * s[:rank].unsqueeze(0)) @ vh[:rank, :] + + +def test_lora_xs_factors_use_principal_components(): + torch.manual_seed(0) + weight = torch.randn(12, 8) + linear_in, linear_out = lora_xs_init_worker.lora_xs_factors(weight, rank=4) + + torch.testing.assert_close(linear_out @ linear_in, _principal(weight, 4)) + assert linear_in.is_contiguous() + assert linear_out.is_contiguous() + + +def test_lora_xs_uses_unfused_forward_for_square_core(monkeypatch): + adapter = _Adapter() + fused = lora_xs.TEFusedLoRALinear(nn.Identity(), adapter) + monkeypatch.setattr(lora_xs.LoRA, "transform", lambda self, module, name, prefix: fused) + monkeypatch.setattr(lora_xs, "ParallelLinearAdapter", _Adapter) + + transformed = lora_xs.LoRAXS().transform(nn.Identity()) + + assert type(transformed) is LoRALinear + assert isinstance(transformed.adapter.activation, lora_xs.LoRAXSCore) + + +def test_lora_xs_trains_only_square_core(): + torch.manual_seed(1) + adapter = _Adapter() + lora_xs.configure_lora_xs_adapter(adapter) + nn.init.normal_(adapter.linear_in.weight) + nn.init.normal_(adapter.linear_out.weight) + + output = adapter(torch.randn(3, 8)).sum() + output.backward() + + assert adapter.linear_in.weight.grad is None + assert adapter.linear_out.weight.grad is None + assert adapter.activation.weight.grad.count_nonzero() > 0 + assert sum(parameter.numel() for parameter in adapter.parameters() if parameter.requires_grad) == adapter.dim**2 + assert output != 0 + assert adapter.activation.weight.abs().max() < 1e-3 + + +@pytest.mark.parametrize("export_raises", [False, True]) +def test_lora_xs_export_materializes_and_restores_standard_lora(monkeypatch, export_raises): + torch.manual_seed(2) + adapter = _Adapter() + lora_xs.configure_lora_xs_adapter(adapter) + nn.init.normal_(adapter.linear_out.weight) + nn.init.normal_(adapter.activation.weight) + wrapped = LoRALinear(nn.Identity(), adapter) + original_out = adapter.linear_out.weight.detach().clone() + original_core = adapter.activation + expected_out = original_out @ original_core.weight + monkeypatch.setattr(lora_xs, "ParallelLinearAdapter", _Adapter) + + if export_raises: + with pytest.raises(RuntimeError), lora_xs.LoRAXS().export_context(wrapped): + torch.testing.assert_close(adapter.linear_out.weight, expected_out) + assert isinstance(adapter.activation, nn.Identity) + raise RuntimeError("export failed") + else: + with lora_xs.LoRAXS().export_context(wrapped): + torch.testing.assert_close(adapter.linear_out.weight, expected_out) + assert isinstance(adapter.activation, nn.Identity) + + torch.testing.assert_close(adapter.linear_out.weight, original_out) + assert adapter.activation is original_core + + +@pytest.mark.parametrize("input_is_parallel", [False, True], ids=["column_parallel", "row_parallel"]) +@pytest.mark.parametrize("tp_rank", [0, 1]) +def test_lora_xs_initialization_reshards_frozen_factors( + monkeypatch, + input_is_parallel, + tp_rank, +): + torch.manual_seed(3) + full_weight = torch.randn(12, 8) + rank = 4 + tp_size = 2 + base_shard_dim = 1 if input_is_parallel else 0 + linear_in_shard_dim = 1 if input_is_parallel else 0 + base_weight = full_weight.chunk(tp_size, dim=base_shard_dim)[tp_rank].clone() + linear_in_shape = [rank, full_weight.shape[1]] + linear_in_shape[linear_in_shard_dim] //= tp_size + core = lora_xs.LoRAXSCore(rank, rank, bias=False) + nn.init.normal_(core.weight) + adapter = SimpleNamespace( + input_is_parallel=input_is_parallel, + dim=rank, + alpha=rank, + linear_in=SimpleNamespace(weight=nn.Parameter(torch.empty(linear_in_shape), requires_grad=False)), + linear_out=SimpleNamespace( + weight=nn.Parameter(torch.empty(full_weight.shape[0] // tp_size, rank), requires_grad=False) + ), + activation=core, + ) + base_linear = SimpleNamespace(weight=nn.Parameter(base_weight)) + monkeypatch.setattr(lora_xs_init_worker, "_all_gather", lambda *args: full_weight) + monkeypatch.setattr( + lora_xs_init_worker, + "_synchronized_lora_xs_factors", + lambda weight, rank, group: lora_xs_init_worker.lora_xs_factors(weight, rank), + ) + + lora_xs_init_worker._init_one_lora_xs_adapter( + base_linear, + adapter, + tp_size, + tp_rank, + None, + ) + + linear_in, linear_out = lora_xs_init_worker.lora_xs_factors(full_weight, rank) + torch.testing.assert_close(base_linear.weight, full_weight.chunk(tp_size, dim=base_shard_dim)[tp_rank]) + torch.testing.assert_close( + adapter.linear_in.weight, + linear_in.chunk(tp_size, dim=linear_in_shard_dim)[tp_rank], + ) + torch.testing.assert_close(adapter.linear_out.weight, linear_out.chunk(tp_size, dim=0)[tp_rank]) + expected_core = torch.zeros(rank, rank) + torch.testing.assert_close(adapter.activation.weight, expected_core, atol=1e-4, rtol=0) + assert not torch.equal(adapter.activation.weight, expected_core) diff --git a/tests/tinker/skyrl_train/test_sample_session_routing.py b/tests/tinker/skyrl_train/test_sample_session_routing.py index be420ede22..1280388e49 100644 --- a/tests/tinker/skyrl_train/test_sample_session_routing.py +++ b/tests/tinker/skyrl_train/test_sample_session_routing.py @@ -17,7 +17,11 @@ skyrl_train_backend = pytest.importorskip("skyrl.backends.skyrl_train_backend") from skyrl.tinker import types # noqa: E402 +from skyrl.tinker.config import EngineConfig # noqa: E402 from skyrl.tinker.engine import prepare_sample_batch # noqa: E402 +from skyrl.tinker.extra.skyrl_train_inference_forwarding import ( # noqa: E402 + _resolve_forwarded_model_name, +) BASE_MODEL = "trl-internal-testing/tiny-Qwen3ForCausalLM" @@ -50,6 +54,7 @@ def _sample_input(**kwargs) -> types.SampleInput: def test_sample_with_remote_client_sets_session_id(monkeypatch): """Test that session_id is set correctly if routing key is present""" + monkeypatch.setattr(skyrl_train_backend, "_uses_lora_weight_sync", lambda cfg: False) monkeypatch.setattr(skyrl_train_backend, "resolve_policy_model_name", lambda cfg: BASE_MODEL) spy = _SpyClient() @@ -74,3 +79,38 @@ def test_sample_with_remote_client_sets_session_id(monkeypatch): sample(fake_self, batch_without_session) assert len(spy.payloads) == 1 assert "session_id" not in spy.payloads[0]["json"] + + +@pytest.mark.parametrize( + ("uses_lora_weight_sync", "expected_model"), + [(True, "model_test"), (False, BASE_MODEL)], +) +def test_sample_with_remote_client_routes_model(monkeypatch, uses_lora_weight_sync, expected_model): + monkeypatch.setattr(skyrl_train_backend, "_uses_lora_weight_sync", lambda cfg: uses_lora_weight_sync) + monkeypatch.setattr(skyrl_train_backend, "resolve_policy_model_name", lambda cfg: BASE_MODEL) + + spy = _SpyClient() + fake_self = SimpleNamespace( + _cfg=None, + _model_ids_to_role={"model_test": "policy"}, + _inference_engine_client=spy, + _aggregate_sample_results=lambda prepared_batch, outputs: {}, + ) + sample = skyrl_train_backend.SkyRLTrainBackend._sample_with_remote_client + sample(fake_self, prepare_sample_batch({"req": ("model_test", _sample_input())})) + + assert spy.payloads[0]["json"]["model"] == expected_model + + +@pytest.mark.parametrize( + ("merge_lora", "expected_model"), + [(False, "model_test"), (True, BASE_MODEL)], +) +def test_forwarded_sample_routes_merged_model(merge_lora, expected_model): + engine_config = EngineConfig( + base_model=BASE_MODEL, + backend="megatron", + backend_config={"trainer.policy.megatron_config.lora_config.merge_lora": merge_lora}, + ) + + assert _resolve_forwarded_model_name(engine_config, "model_test", None) == expected_model