From c58591fb633b74ea83f7ee75c5a528347f7a194b Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Wed, 12 Aug 2026 10:53:20 +0800 Subject: [PATCH 1/9] feat: add zero-shot spectral LoRA allocation --- cookbook/transformers/zero_shot_lora.py | 175 ++++++++++++ cookbook/transformers/zero_shot_lora.sh | 30 ++ scripts/compute_zero_shot_hybrid_config.py | 107 +++++++ .../model/transformers/strategy/accelerate.py | 10 +- .../transformers/strategy/native_fsdp.py | 9 +- .../model/transformers/transformers.py | 8 +- .../model/transformers/zero_shot_lora.py | 265 ++++++++++++++++++ tests/transformers/test_zero_shot_lora.py | 244 ++++++++++++++++ 8 files changed, 842 insertions(+), 6 deletions(-) create mode 100644 cookbook/transformers/zero_shot_lora.py create mode 100644 cookbook/transformers/zero_shot_lora.sh create mode 100644 scripts/compute_zero_shot_hybrid_config.py create mode 100644 src/twinkle/model/transformers/zero_shot_lora.py create mode 100644 tests/transformers/test_zero_shot_lora.py diff --git a/cookbook/transformers/zero_shot_lora.py b/cookbook/transformers/zero_shot_lora.py new file mode 100644 index 000000000..710b56c17 --- /dev/null +++ b/cookbook/transformers/zero_shot_lora.py @@ -0,0 +1,175 @@ +import json +from pathlib import Path + +from peft import LoraConfig + +import twinkle +from twinkle import DeviceMesh, get_device_placement, get_logger +from twinkle.cli import CLI +from twinkle.dataloader import DataLoader +from twinkle.dataset import Dataset, DatasetMeta +from twinkle.model import TransformersModel +from twinkle.model.transformers.zero_shot_lora import ( + allocate_zero_shot_modules, + build_zero_shot_lora_config, + build_zero_shot_param_groups, + compute_zero_shot_scores, + select_zero_shot_targets, +) + +logger = get_logger() +args = CLI.from_args() + +device_mesh = DeviceMesh.from_sizes(fsdp_size=args.infra.fsdp_size, dp_size=args.infra.dp_size) +twinkle.initialize(mode=args.infra.mode, global_device_mesh=device_mesh) + + +def build_dataset(data_slice) -> Dataset: + dataset = Dataset(dataset_meta=DatasetMeta(args.dataset.dataset_id, data_slice=data_slice)) + dataset.set_template( + args.template.template_cls, + model_id=args.model.model_id, + max_length=args.template.max_length, + truncation_strategy=args.template.truncation_strategy, + enable_thinking=args.template.enable_thinking, + ) + dataset.encode(num_proc=8, load_from_cache_file=True) + return dataset + + +def load_zero_shot_config(config_value: str) -> LoraConfig: + """Load an allocation produced by scripts/compute_zero_shot_hybrid_config.py.""" + config_path = Path(config_value).expanduser() + if not config_path.is_file(): + raise FileNotFoundError(f'Zero-shot config JSON file not found: {config_path}') + with config_path.open(encoding='utf-8') as handle: + raw_config = json.load(handle) + if not isinstance(raw_config, dict): + raise ValueError('Zero-shot config JSON must contain an object.') + if raw_config.get('method') not in (None, 'zero_shot_spectral'): + raise ValueError(f'Unsupported zero-shot config method: {raw_config.get("method")!r}.') + + def module_list(primary_key, peft_key): + value = raw_config.get(primary_key, raw_config.get(peft_key)) + if value is None: + return [] + if isinstance(value, str): + value = [value] + if not isinstance(value, (list, tuple, set)) or not all(isinstance(item, str) for item in value): + raise ValueError(f'Zero-shot config {primary_key} must be a list of module names.') + return sorted(set(value)) + + s_fft = module_list('s_fft', 'modules_to_save') + s_lora = module_list('s_lora', 'target_modules') + overlap = set(s_fft) & set(s_lora) + if overlap: + raise ValueError(f'Zero-shot modules cannot be both FFT and LoRA: {sorted(overlap)}') + r = int(raw_config.get('r', args.extra.get('zero_shot_r', args.lora.lora_r))) + lora_alpha = int(raw_config.get('lora_alpha', args.extra.get('zero_shot_alpha', r * 2))) + lora_dropout = float(raw_config.get('lora_dropout', 0.0)) + return build_zero_shot_lora_config( + s_lora, + s_fft, + r=r, + lora_alpha=lora_alpha, + lora_dropout=lora_dropout, + ) + + +def compute_allocation() -> LoraConfig: + """Allocate full fine-tuning and LoRA modules from pretrained spectra.""" + r = int(args.extra.get('zero_shot_r', args.lora.lora_r)) + lora_alpha = int(args.extra.get('zero_shot_alpha', r * 2)) + fft_ratio = float(args.extra.get('zero_shot_fft_ratio', 0.1)) + epsilon = float(args.extra.get('zero_shot_epsilon', 1e-12)) + + logger.info(f'Zero-shot allocation: loading pretrained model ' + f'(r={r}, alpha={lora_alpha}, FFT budget={fft_ratio:.1%})') + base_model = TransformersModel(model_id=args.model.model_id) + if base_model._memory_efficient_init: + raise ValueError('Zero-shot spectral scoring requires materialized weights; ' + 'disable memory_efficient_init.') + target_config = LoraConfig(**args.get_lora_args()) + targets = select_zero_shot_targets(base_model.model, target_config) + param_counts = {name: module.weight.numel() for name, module in targets.items()} + cache_dir = Path(args.training.output_dir) / 'zero-shot-spectrum-cache' + scores = compute_zero_shot_scores( + base_model.model, + target_config, + r=r, + cache_dir=cache_dir, + cache_key=str(args.model.model_id), + epsilon=epsilon, + log_interval=args.training.log_interval, + ) + s_fft, s_lora = allocate_zero_shot_modules(scores, param_counts, fft_ratio=fft_ratio) + ranked = sorted(scores, key=lambda name: (-scores[name], name)) + logger.info('Zero-shot top-10 modules: ' + ', '.join( + f'{name}=score:{scores[name]:.4f},effective_rank:{scores.metrics[name]["effective_rank"]:.1f},' + f'coverage:{scores.metrics[name]["rank_coverage"]:.4f},' + f'condition:{scores.metrics[name]["condition_number"]:.2e},' + f'decay:{scores.metrics[name]["decay"]:.4e}' + for name in ranked[:10])) + fft_params = sum(param_counts[name] for name in s_fft) + total_params = sum(param_counts.values()) + logger.info(f'Zero-shot allocation: {len(s_fft)} FFT modules, {len(s_lora)} LoRA modules ' + f'({fft_params / total_params:.1%} of candidate params to FFT; cache={cache_dir})') + logger.info(f'Zero-shot FFT modules: {", ".join(s_fft) if s_fft else "(none)"}') + del base_model + return build_zero_shot_lora_config(s_lora, s_fft, r=r, lora_alpha=lora_alpha) + + +def train() -> None: + train_samples = args.training.train_samples or 1000 + config_value = args.extra.get('zero_shot_config') + if config_value: + zero_shot_config = load_zero_shot_config(config_value) + logger.info(f'Using supplied zero-shot config; skipping spectral scoring ' + f'({len(zero_shot_config.modules_to_save or [])} FFT modules, ' + f'{len(zero_shot_config.target_modules or [])} LoRA modules)') + else: + zero_shot_config = compute_allocation() + + dataset = build_dataset(range(train_samples)) + dataloader = DataLoader(dataset=dataset, batch_size=args.training.batch_size) + model = TransformersModel(model_id=args.model.model_id) + model.add_adapter_to_model( + args.lora.adapter_name, + zero_shot_config, + gradient_accumulation_steps=args.training.gradient_accumulation_steps, + ) + param_groups = build_zero_shot_param_groups( + model.strategy.unwrap_model(model.model), + lr_lora=float(args.extra.get('zero_shot_lr_lora', args.optimizer.learning_rate)), + lr_fft=float(args.extra.get('zero_shot_lr_fft', 1e-6)), + weight_decay=args.optimizer.weight_decay, + adapter_name=args.lora.adapter_name, + ) + model.set_optimizer(optimizer_cls=args.optimizer.optimizer_cls, params=param_groups) + model.set_lr_scheduler( + scheduler_cls=args.scheduler.scheduler_cls, + num_warmup_steps=args.scheduler.num_warmup_steps, + num_training_steps=len(dataloader), + ) + + logger.info(get_device_placement()) + logger.info(model.get_train_configs()) + optimizer_group = model.optimizer_group[args.lora.adapter_name] + for batch in dataloader: + model.forward_backward(inputs=batch) + model.clip_grad_and_step() + cur_step = optimizer_group.cur_step + if cur_step % args.training.log_interval == 0: + logger.info(f'step {cur_step}/{len(dataloader)}, ' + f'metric: {model.calculate_metric(is_training=True)}') + model.save( + 'last-checkpoint', + output_dir=args.training.output_dir, + adapter_name=args.lora.adapter_name, + save_optimizer=True, + consumed_train_samples=dataloader.get_state()['consumed_train_samples'], + ) + + +if __name__ == '__main__': + train() diff --git a/cookbook/transformers/zero_shot_lora.sh b/cookbook/transformers/zero_shot_lora.sh new file mode 100644 index 000000000..a5b9dfd76 --- /dev/null +++ b/cookbook/transformers/zero_shot_lora.sh @@ -0,0 +1,30 @@ +#!/bin/sh +# Generate the allocation from pretrained spectra, then train the selected FFT/LoRA modules. +# Reuse an offline allocation with: --zero-shot-config ./output/zero_shot/hybrid_config.json + +CUDA_VISIBLE_DEVICES=0,1,2,3 \ + torchrun --nproc_per_node=4 zero_shot_lora.py \ + --model-id ms://Qwen/Qwen3.5-9B \ + --dataset-id data/financial_sft/processed/finqa_tatqa_train_messages.jsonl \ + --template-cls Qwen3_5Template \ + --fsdp-size 4 \ + --dp-size 1 \ + --batch-size 16 \ + --optimizer-cls AdamW \ + --lr 2.5e-5 \ + --weight-decay 0.01 \ + --gradient-accumulation-steps 4 \ + --log-interval 1 \ + --output-dir ./output/zero_shot_lora \ + --adapter-name default \ + --lora-r 64 \ + --scheduler-cls CosineWarmupScheduler \ + --num-warmup-steps 10 \ + --train-samples 1000 \ + --zero-shot-r 64 \ + --zero-shot-alpha 128 \ + --zero-shot-fft-ratio 0.3 \ + --zero-shot-epsilon 1e-12 \ + --zero-shot-lr-fft 1e-6 \ + --zero-shot-lr-lora 2.5e-5 \ + "$@" diff --git a/scripts/compute_zero_shot_hybrid_config.py b/scripts/compute_zero_shot_hybrid_config.py new file mode 100644 index 000000000..12fd60cdf --- /dev/null +++ b/scripts/compute_zero_shot_hybrid_config.py @@ -0,0 +1,107 @@ +#!/usr/bin/env python3 +"""Compute a zero-shot FFT/LoRA allocation from pretrained weight spectra. + +Example: + python scripts/compute_zero_shot_hybrid_config.py \ + --model-id ms://Qwen/Qwen3.5-9B \ + --zero-shot-r 64 \ + --zero-shot-fft-ratio 0.3 \ + --output-dir ./output/zero_shot \ + --hybrid-config-output ./output/zero_shot/hybrid_config.json +""" + +import json +import os +from pathlib import Path +from peft import LoraConfig + +import twinkle +from twinkle import DeviceMesh, Platform, get_logger +from twinkle.cli import CLI +from twinkle.model import TransformersModel +from twinkle.model.transformers.zero_shot_lora import (CANDIDATE_TYPES, allocate_zero_shot_modules, + compute_zero_shot_scores, select_zero_shot_targets) + +logger = get_logger() +args = CLI.from_args() + +# This utility is intentionally single-process: it only reads weights and writes one JSON config. +device_mesh = DeviceMesh.from_sizes(fsdp_size=1, dp_size=1) +twinkle.initialize(mode=args.infra.mode, global_device_mesh=device_mesh) + + +def main() -> None: + if not args.model.model_id: + raise ValueError('--model-id is required.') + + r = int(args.extra.get('zero_shot_r', args.lora.lora_r)) + lora_alpha = int(args.extra.get('zero_shot_alpha', r * 2)) + fft_ratio = float(args.extra.get('zero_shot_fft_ratio', 0.1)) + epsilon = float(args.extra.get('zero_shot_epsilon', 1e-12)) + output_path = Path(args.extra.get('hybrid_config_output', + Path(args.training.output_dir) / 'hybrid_config.json')).expanduser() + cache_dir = Path( + args.extra.get('zero_shot_cache_dir', + Path(args.training.output_dir) / 'zero-shot-spectrum-cache')).expanduser() + + logger.info(f'Loading pretrained model for zero-shot spectral scoring: {args.model.model_id}') + model = TransformersModel(model_id=args.model.model_id) + if model._memory_efficient_init: + raise ValueError('Zero-shot spectral scoring requires materialized weights; disable memory_efficient_init.') + + target_config = LoraConfig( + r=r, + lora_alpha=lora_alpha, + lora_dropout=0.0, + target_modules=list(CANDIDATE_TYPES.values()), + ) + targets = select_zero_shot_targets(model.model, target_config) + param_counts = {name: module.weight.numel() for name, module in targets.items()} + scores = compute_zero_shot_scores( + model.model, + target_config, + r=r, + cache_dir=cache_dir, + cache_key=str(args.model.model_id), + epsilon=epsilon, + log_interval=args.training.log_interval, + ) + counts = {name: param_counts[name] for name in scores} + s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=fft_ratio) + fft_params = sum(counts[name] for name in s_fft) + total_params = sum(counts.values()) + + config = { + 'method': 'zero_shot_spectral', + 'model_id': args.model.model_id, + 's_fft': s_fft, + 's_lora': s_lora, + 'r': r, + 'lora_alpha': lora_alpha, + 'lora_dropout': 0.0, + 'fft_ratio': fft_ratio, + 'realized_fft_param_ratio': fft_params / total_params, + 'zero_shot_epsilon': epsilon, + 'metrics': { + name: { + 'score': scores[name], + **scores.metrics[name], + } + for name in sorted(scores) + }, + } + + if Platform.is_master(): + output_path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = output_path.with_suffix(f'{output_path.suffix}.tmp') + with temporary_path.open('w', encoding='utf-8') as handle: + json.dump(config, handle, ensure_ascii=False, indent=2, sort_keys=True) + handle.write('\n') + os.replace(temporary_path, output_path) + logger.info(f'Zero-shot config written to {output_path}: ' + f'{len(s_fft)} FFT modules, {len(s_lora)} LoRA modules, ' + f'realized FFT parameter ratio={fft_params / total_params:.2%}') + + +if __name__ == '__main__': + main() diff --git a/src/twinkle/model/transformers/strategy/accelerate.py b/src/twinkle/model/transformers/strategy/accelerate.py index 3bf627e9d..fdd808355 100644 --- a/src/twinkle/model/transformers/strategy/accelerate.py +++ b/src/twinkle/model/transformers/strategy/accelerate.py @@ -6,6 +6,8 @@ from twinkle import DeviceMesh from .load_context import fsdp_pretrained_load_context +MODULES_TO_SAVE_SEGMENT = '.modules_to_save.' + class AccelerateStrategy: """A training strategy that uses `accelerate` to wrap models. @@ -204,13 +206,13 @@ def get_full_state_dict(self, model) -> dict: return state_dict def get_adapter_state_dict(self, model, adapter_name: str) -> dict: - """Collect only LoRA adapter parameters.""" + """Collect LoRA parameters and fully trained PEFT modules.""" from twinkle.utils import torch_util unwrapped = self.unwrap_model(model) state_dict = {} adapter_suffix = f'.{adapter_name}.' for name, param in unwrapped.named_parameters(): - if not _is_lora_state_key(name) or adapter_suffix not in name: + if not _is_saveable_adapter_key(name) or adapter_suffix not in name: continue local = torch_util.to_local_tensor(param) state_dict[name] = local.cpu() @@ -220,3 +222,7 @@ def get_adapter_state_dict(self, model, adapter_name: str) -> dict: def _is_lora_state_key(name: str) -> bool: return 'lora_A' in name or 'lora_B' in name or 'lora_embedding' in name + + +def _is_saveable_adapter_key(name: str) -> bool: + return _is_lora_state_key(name) or MODULES_TO_SAVE_SEGMENT in name diff --git a/src/twinkle/model/transformers/strategy/native_fsdp.py b/src/twinkle/model/transformers/strategy/native_fsdp.py index 2722a900d..150026920 100644 --- a/src/twinkle/model/transformers/strategy/native_fsdp.py +++ b/src/twinkle/model/transformers/strategy/native_fsdp.py @@ -17,6 +17,7 @@ logger = get_logger() LORA_STATE_KEY_MARKERS = ('lora_A', 'lora_B', 'lora_embedding') +MODULES_TO_SAVE_SEGMENT = '.modules_to_save.' PEFT_BASE_PREFIX = 'base_model.model.' PEFT_BASE_LAYER_SEGMENT = 'base_layer' @@ -346,7 +347,7 @@ def get_full_state_dict(self, model) -> dict: return state_dict def get_adapter_state_dict(self, model, adapter_name: str) -> dict: - """Collect only LoRA adapter parameters, with EP-aware all-gather.""" + """Collect LoRA parameters and fully trained PEFT modules, with EP-aware all-gather.""" unwrapped = self.unwrap_model(model) state_dict = {} @@ -361,7 +362,7 @@ def get_adapter_state_dict(self, model, adapter_name: str) -> dict: adapter_suffix = f'.{adapter_name}.' for name, param in unwrapped.named_parameters(): - if not _is_lora_state_key(name) or adapter_suffix not in name: + if not _is_saveable_adapter_key(name) or adapter_suffix not in name: continue local_full = torch_util.to_local_tensor(param) @@ -694,6 +695,10 @@ def _is_lora_state_key(name: str) -> bool: return any(marker in name for marker in LORA_STATE_KEY_MARKERS) +def _is_saveable_adapter_key(name: str) -> bool: + return _is_lora_state_key(name) or MODULES_TO_SAVE_SEGMENT in name + + def _strip_peft_base_prefix(name: str) -> str: while name.startswith(PEFT_BASE_PREFIX): name = name[len(PEFT_BASE_PREFIX):] diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index c8fec86da..9a568452a 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -956,7 +956,7 @@ def set_optimizer(self, optimizer_cls: Union[Type[Optimizer], str, Optimizer], * def _get_trainable_parameters(self, adapter_name=_default_adapter_name): is_default = adapter_name == _default_adapter_name - pattern = re.compile(rf'\.lora_\w+\.{re.escape(adapter_name)}\.') + pattern = re.compile(rf'\.(?:lora_\w+|modules_to_save)\.{re.escape(adapter_name)}\.') params = {} model = self.strategy.unwrap_model(self.model) for name, param in model.named_parameters(): @@ -1044,9 +1044,13 @@ def _get_adapter_state_dict_for_save(self, adapter_name: str) -> dict: # Avoid collecting the full base model for large FSDP/EP jobs. adapter_state = self.strategy.get_adapter_state_dict(self.model, adapter_name) adapter_suffix = f'.{adapter_name}.' + modules_to_save_infix = f'.modules_to_save.{adapter_name}.' processed_state_dict = {} for key, value in adapter_state.items(): - normalized = key.replace(adapter_suffix, '.') + if modules_to_save_infix in key: + normalized = key.replace(modules_to_save_infix, '.') + else: + normalized = key.replace(adapter_suffix, '.') processed_state_dict[normalized] = value return processed_state_dict diff --git a/src/twinkle/model/transformers/zero_shot_lora.py b/src/twinkle/model/transformers/zero_shot_lora.py new file mode 100644 index 000000000..3c9cf5741 --- /dev/null +++ b/src/twinkle/model/transformers/zero_shot_lora.py @@ -0,0 +1,265 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import copy +import hashlib +import os +import re +import torch +import torch.distributed as dist +from pathlib import Path +from peft import LoraConfig +from peft.tuners.tuners_utils import _maybe_include_all_linear_layers, check_target_module_exists +from torch import nn +from typing import Dict, List, Mapping, Optional, Tuple + +from twinkle import get_logger + +logger = get_logger() + +CANDIDATE_TYPES: Dict[str, str] = { + 'q': 'q_proj', + 'k': 'k_proj', + 'v': 'v_proj', + 'o': 'o_proj', + 'gate': 'gate_proj', + 'up': 'up_proj', + 'down': 'down_proj', +} +_CANDIDATE_SUFFIXES = frozenset(CANDIDATE_TYPES.values()) +_LAYER_RE = re.compile(r'\blayers\.(\d+)\.') + + +class ZeroShotScores(dict[str, float]): + """Zero-shot spectral scores with per-module metric details.""" + + def __init__(self, scores: Mapping[str, float], metrics: Mapping[str, Mapping[str, float]]) -> None: + super().__init__(scores) + self.metrics = {name: dict(values) for name, values in metrics.items()} + + +def compute_zero_shot_spectral_metrics( + singular_values: torch.Tensor, + r: int, + epsilon: float = 1e-12, +) -> Dict[str, float]: + """Compute the pretrained-spectrum metrics used for zero-shot allocation.""" + if r <= 0: + raise ValueError('Zero-shot LoRA rank r must be positive.') + if singular_values.ndim != 1 or singular_values.numel() == 0: + raise ValueError('Zero-shot scoring requires a non-empty singular-value vector.') + if epsilon <= 0: + raise ValueError('Zero-shot scoring epsilon must be positive.') + + values = singular_values.detach().to(device='cpu', dtype=torch.float64).abs() + values = values.sort(descending=True).values + stable_values = values.clamp_min(epsilon) + probabilities = stable_values / stable_values.sum() + entropy = -(probabilities * probabilities.log()).sum() + effective_rank = entropy.exp() + dimension = values.numel() + + rank = min(r, dimension) + squared = values.square() + rank_coverage = squared[:rank].sum() / squared.sum().clamp_min(epsilon) + condition_number = values[0] / values[-1].clamp_min(epsilon) + if dimension > 1: + decay = torch.diff(stable_values.log()).abs().mean() + inverse_decay = 1.0 / decay.clamp_min(epsilon) + else: + decay = torch.tensor(float('inf'), dtype=torch.float64) + inverse_decay = torch.tensor(0.0, dtype=torch.float64) + + score = (0.3 * (effective_rank / dimension) + 0.3 * (1.0 - rank_coverage) + 0.2 * + (torch.log1p(condition_number) / 10.0) + 0.2 * inverse_decay) + return { + 'effective_rank': float(effective_rank.item()), + 'normalized_effective_rank': float((effective_rank / dimension).item()), + 'rank_coverage': float(rank_coverage.item()), + 'condition_number': float(condition_number.item()), + 'decay': float(decay.item()), + 'score': float(score.item()), + } + + +def _is_candidate_module(module_name: str) -> bool: + return _LAYER_RE.search(module_name) is not None and module_name.rsplit('.', 1)[-1] in _CANDIDATE_SUFFIXES + + +def select_zero_shot_targets(model: nn.Module, config: LoraConfig) -> Dict[str, nn.Module]: + """Resolve materialized linear modules eligible for zero-shot allocation.""" + resolved = _maybe_include_all_linear_layers(copy.deepcopy(config), model) + targets: Dict[str, nn.Module] = {} + for name, module in model.named_modules(): + weight = getattr(module, 'weight', None) + if not name or not isinstance(weight, nn.Parameter) or weight.ndim != 2: + continue + if not check_target_module_exists(resolved, name) or not _is_candidate_module(name): + continue + if not isinstance(module, nn.Linear): + raise ValueError(f'Zero-shot LoRA target {name!r} is not an nn.Linear module.') + if weight.is_meta: + raise ValueError('Zero-shot spectral scoring requires materialized weights; ' + 'disable memory_efficient_init.') + targets[name] = module + if not targets: + raise ValueError(f'Zero-shot LoRA found no candidate modules for {config.target_modules!r}.') + return targets + + +def _spectrum_cache_path( + cache_dir: Path, + cache_key: str, + module_name: str, + shape: Tuple[int, int], + dtype: torch.dtype, +) -> Path: + identity = f'{cache_key}|{module_name}|{shape[0]}x{shape[1]}|{dtype}' + digest = hashlib.sha256(identity.encode('utf-8')).hexdigest()[:24] + return cache_dir / f'spectrum-{digest}.pt' + + +@torch.no_grad() +def compute_zero_shot_scores( + model: nn.Module, + config: LoraConfig, + r: int, + *, + cache_dir: Optional[Path] = None, + cache_key: str = '', + epsilon: float = 1e-12, + log_interval: int = 20, +) -> ZeroShotScores: + """Score pretrained modules from singular-value spectra without training data.""" + targets = select_zero_shot_targets(model, config) + names = sorted(targets) + distributed = dist.is_available() and dist.is_initialized() + rank = dist.get_rank() if distributed else 0 + scores: Dict[str, float] = {} + metrics_by_module: Dict[str, Dict[str, float]] = {} + + if rank == 0: + if log_interval > 0: + logger.info(f'Zero-shot spectral scoring: {len(names)} modules, LoRA rank={r}') + for index, name in enumerate(names, start=1): + weight = targets[name].weight.detach() + shape = tuple(weight.shape) + spectrum_path = None + if cache_dir is not None: + spectrum_path = _spectrum_cache_path(Path(cache_dir), cache_key, name, shape, weight.dtype) + + singular_values = None + if spectrum_path is not None and spectrum_path.is_file(): + try: + cached = torch.load(spectrum_path, map_location='cpu', weights_only=True) + if isinstance(cached, torch.Tensor) and cached.ndim == 1 and cached.numel() == min(shape): + singular_values = cached + except (OSError, RuntimeError, EOFError): + logger.warning(f'Zero-shot spectrum cache is unreadable: {spectrum_path}; recomputing') + + if singular_values is None: + cpu_weight = weight.to(device='cpu') + if cpu_weight.dtype not in (torch.float32, torch.float64): + cpu_weight = cpu_weight.float() + singular_values = torch.linalg.svdvals(cpu_weight) + if spectrum_path is not None: + spectrum_path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = spectrum_path.with_suffix('.tmp') + torch.save(singular_values, temporary_path) + os.replace(temporary_path, spectrum_path) + + metrics = compute_zero_shot_spectral_metrics(singular_values, r=r, epsilon=epsilon) + metrics_by_module[name] = metrics + scores[name] = metrics['score'] + if log_interval > 0 and (index % log_interval == 0 or index == len(names)): + logger.info(f'Zero-shot spectral scoring: {index}/{len(names)} modules ' + f'(last {name} -> score={metrics["score"]:.4f}, ' + f'effective_rank={metrics["effective_rank"]:.1f}, ' + f'rank_coverage={metrics["rank_coverage"]:.4f})') + + if distributed: + payload = [(scores, metrics_by_module) if rank == 0 else None] + dist.broadcast_object_list(payload, src=0) + scores, metrics_by_module = payload[0] + return ZeroShotScores(scores, metrics_by_module) + + +def allocate_zero_shot_modules( + scores: Mapping[str, float], + param_counts: Mapping[str, int], + fft_ratio: float = 0.1, +) -> Tuple[List[str], List[str]]: + """Allocate the highest-scoring prefix to full fine-tuning within a parameter budget.""" + if not 0.0 <= fft_ratio < 1.0: + raise ValueError('Zero-shot FFT ratio must be in [0, 1).') + missing = set(scores) - set(param_counts) + if missing: + raise ValueError(f'Zero-shot allocation is missing parameter counts for: {sorted(missing)}.') + total = sum(param_counts[name] for name in scores) + if total <= 0: + raise ValueError('Zero-shot candidate parameter count must be positive.') + + budget = fft_ratio * total + ordered = sorted(scores, key=lambda name: (-scores[name], name)) + s_fft: List[str] = [] + used = 0 + for name in ordered: + cost = param_counts[name] + if used + cost > budget + 1e-9: + break + s_fft.append(name) + used += cost + s_fft_set = set(s_fft) + s_lora = [name for name in ordered if name not in s_fft_set] + return sorted(s_fft), sorted(s_lora) + + +def build_zero_shot_lora_config( + s_lora: List[str], + s_fft: List[str], + r: int = 16, + lora_alpha: int = 32, + lora_dropout: float = 0.0, + **kwargs, +) -> LoraConfig: + if not s_lora: + raise ValueError('Zero-shot LoRA requires at least one LoRA module.') + return LoraConfig( + r=r, + lora_alpha=lora_alpha, + lora_dropout=lora_dropout, + target_modules=list(s_lora), + modules_to_save=list(s_fft), + **kwargs, + ) + + +def build_zero_shot_param_groups( + peft_model: nn.Module, + lr_lora: float = 2.5e-5, + lr_fft: float = 1e-6, + weight_decay: float = 0.0, + adapter_name: str = 'default', +) -> List[dict]: + """Build separate optimizer groups for LoRA and fully fine-tuned modules.""" + lora_params, lora_names = [], [] + fft_params, fft_names = [], [] + adapter_token = f'.{adapter_name}.' + for name, param in peft_model.named_parameters(): + if not param.requires_grad: + continue + if '.lora_' in name and adapter_token in name: + lora_params.append(param) + lora_names.append(name) + elif '.modules_to_save.' in name and adapter_token in name: + fft_params.append(param) + fft_names.append(name) + else: + raise ValueError(f'Zero-shot LoRA cannot classify trainable parameter {name!r}.') + + groups: List[dict] = [] + if lora_params: + groups.append({'params': lora_params, 'param_names': lora_names, 'lr': lr_lora, 'weight_decay': weight_decay}) + if fft_params: + groups.append({'params': fft_params, 'param_names': fft_names, 'lr': lr_fft, 'weight_decay': weight_decay}) + if not groups: + raise ValueError('Zero-shot LoRA found no trainable parameters to optimize.') + return groups diff --git a/tests/transformers/test_zero_shot_lora.py b/tests/transformers/test_zero_shot_lora.py new file mode 100644 index 000000000..7db4307d9 --- /dev/null +++ b/tests/transformers/test_zero_shot_lora.py @@ -0,0 +1,244 @@ +import pytest +import torch +from peft import LoraConfig, PeftModel, get_peft_model +from peft.utils import get_peft_model_state_dict +from torch import nn + +from twinkle.model.transformers.zero_shot_lora import ( + CANDIDATE_TYPES, + allocate_zero_shot_modules, + build_zero_shot_lora_config, + build_zero_shot_param_groups, + compute_zero_shot_scores, + compute_zero_shot_spectral_metrics, + select_zero_shot_targets, +) + + +class TinyDecoder(nn.Module): + + def __init__(self, num_layers=2, dim=8): + super().__init__() + self.layers = nn.ModuleList() + for _ in range(num_layers): + layer = nn.Module() + layer.self_attn = nn.Module() + layer.self_attn.q_proj = nn.Linear(dim, dim, bias=False) + layer.self_attn.k_proj = nn.Linear(dim, dim, bias=False) + layer.self_attn.v_proj = nn.Linear(dim, dim, bias=False) + layer.self_attn.o_proj = nn.Linear(dim, dim, bias=False) + layer.mlp = nn.Module() + layer.mlp.gate_proj = nn.Linear(dim, dim, bias=False) + layer.mlp.up_proj = nn.Linear(dim, dim, bias=False) + layer.mlp.down_proj = nn.Linear(dim, dim, bias=False) + self.layers.append(layer) + + def forward(self, inputs): + for layer in self.layers: + hidden = layer.self_attn.q_proj(inputs) + inputs = layer.mlp.down_proj(layer.mlp.up_proj(hidden)) + return inputs + + +def test_spectral_metrics_match_weighted_formula(): + singular_values = torch.tensor([4.0, 2.0, 1.0, 0.5], dtype=torch.float64) + metrics = compute_zero_shot_spectral_metrics(singular_values, r=1) + + probabilities = singular_values / singular_values.sum() + effective_rank = torch.exp(-(probabilities * probabilities.log()).sum()).item() + rank_coverage = (singular_values[0].square() / singular_values.square().sum()).item() + condition_number = 8.0 + decay = torch.diff(singular_values.log()).abs().mean().item() + expected = ( + 0.3 * (effective_rank / singular_values.numel()) + + 0.3 * (1.0 - rank_coverage) + + 0.2 * (torch.log1p(torch.tensor(condition_number)).item() / 10.0) + + 0.2 * (1.0 / decay) + ) + + assert metrics['effective_rank'] == pytest.approx(effective_rank) + assert metrics['rank_coverage'] == pytest.approx(rank_coverage) + assert metrics['condition_number'] == pytest.approx(condition_number) + assert metrics['decay'] == pytest.approx(decay) + assert metrics['score'] == pytest.approx(expected) + + +@pytest.mark.parametrize('singular_values,r,message', [ + (torch.ones(4), 0, 'rank r must be positive'), + (torch.empty(0), 1, 'non-empty singular-value vector'), +]) +def test_spectral_metrics_validate_inputs(singular_values, r, message): + with pytest.raises(ValueError, match=message): + compute_zero_shot_spectral_metrics(singular_values, r=r) + + +def test_select_targets_covers_supported_module_types(): + model = TinyDecoder(num_layers=2) + config = LoraConfig(r=4, target_modules=list(CANDIDATE_TYPES.values())) + + targets = select_zero_shot_targets(model, config) + + assert len(targets) == 14 + assert all(name.rsplit('.', 1)[-1] in CANDIDATE_TYPES.values() for name in targets) + + +def test_scores_cache_pretrained_singular_values(tmp_path, monkeypatch): + model = TinyDecoder(num_layers=1) + config = LoraConfig(r=2, target_modules=['q_proj', 'down_proj']) + original_svdvals = torch.linalg.svdvals + calls = [] + + def record_svdvals(weight): + calls.append(tuple(weight.shape)) + return original_svdvals(weight) + + monkeypatch.setattr(torch.linalg, 'svdvals', record_svdvals) + scores = compute_zero_shot_scores( + model, config, r=2, cache_dir=tmp_path, cache_key='tiny', log_interval=0) + cached_scores = compute_zero_shot_scores( + model, config, r=2, cache_dir=tmp_path, cache_key='tiny', log_interval=0) + + assert set(scores) == {'layers.0.mlp.down_proj', 'layers.0.self_attn.q_proj'} + assert scores == cached_scores + assert len(calls) == 2 + assert len(list(tmp_path.glob('spectrum-*.pt'))) == 2 + assert all('effective_rank' in scores.metrics[name] for name in scores) + + +def test_allocation_prioritizes_high_scores_within_budget(): + scores = {'a': 0.2, 'b': 0.8, 'c': 0.4, 'd': 0.6} + counts = {name: 100 for name in scores} + + s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=0.25) + + assert s_fft == ['b'] + assert set(s_lora) == {'a', 'c', 'd'} + + +def test_allocation_uses_strict_ranked_prefix(): + scores = {'big': 0.9, 'small_a': 0.8, 'small_b': 0.7} + counts = {'big': 80, 'small_a': 30, 'small_b': 15} + + s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=0.7) + + assert s_fft == ['big'] + assert set(s_lora) == {'small_a', 'small_b'} + + +def test_allocation_rejects_full_fft_budget(): + with pytest.raises(ValueError, match=r'\[0, 1\)'): + allocate_zero_shot_modules({'module': 1.0}, {'module': 10}, fft_ratio=1.0) + + +def test_config_requires_a_lora_target(): + with pytest.raises(ValueError, match='at least one LoRA module'): + build_zero_shot_lora_config([], ['layers.0.self_attn.q_proj']) + + +def test_config_and_param_groups_cover_every_trainable_parameter(): + config = build_zero_shot_lora_config( + s_lora=['layers.0.mlp.down_proj'], + s_fft=['layers.0.self_attn.q_proj'], + r=4, + lora_alpha=8, + ) + model = get_peft_model(TinyDecoder(num_layers=1), config) + + groups = build_zero_shot_param_groups(model, lr_lora=2.5e-5, lr_fft=1e-6) + + assert {group['lr'] for group in groups} == {2.5e-5, 1e-6} + grouped = {id(param) for group in groups for param in group['params']} + trainable = {id(param) for param in model.parameters() if param.requires_grad} + assert grouped == trainable + + +@pytest.mark.parametrize('strategy_cls', [ + pytest.param('accelerate', id='accelerate'), + pytest.param('native_fsdp', id='native-fsdp'), +]) +def test_strategy_adapter_state_includes_full_modules(strategy_cls): + if strategy_cls == 'accelerate': + from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy as Strategy + strategy = object.__new__(Strategy) + strategy.unwrap_model = lambda model, *args: model + else: + from twinkle.model.transformers.strategy.native_fsdp import NativeFSDPStrategy as Strategy + strategy = object.__new__(Strategy) + strategy.ep_fsdp_device_mesh = None + + config = build_zero_shot_lora_config( + s_lora=['layers.0.mlp.down_proj'], + s_fft=['layers.0.self_attn.q_proj'], + r=4, + lora_alpha=8, + ) + model = get_peft_model(TinyDecoder(num_layers=1), config) + + state = strategy.get_adapter_state_dict(model, 'default') + + assert any('.lora_A.' in name for name in state) + assert any('.modules_to_save.default.' in name for name in state) + + +def test_twinkle_checkpoint_normalization_round_trips_full_modules(tmp_path): + from safetensors.torch import save_file + from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy + from twinkle.model.transformers.transformers import TransformersModel + + torch.manual_seed(0) + base = TinyDecoder(num_layers=1) + base_state = {name: value.detach().clone() for name, value in base.state_dict().items()} + config = build_zero_shot_lora_config( + s_lora=['layers.0.mlp.down_proj'], + s_fft=['layers.0.self_attn.q_proj'], + r=4, + lora_alpha=8, + ) + peft_model = get_peft_model(base, config) + with torch.no_grad(): + for param in peft_model.parameters(): + if param.requires_grad: + param.add_(0.05) + inputs = torch.randn(2, 8) + expected = peft_model(inputs).detach() + + strategy = object.__new__(AccelerateStrategy) + strategy.unwrap_model = lambda model, *args: model + twinkle_model = object.__new__(TransformersModel) + twinkle_model.strategy = strategy + twinkle_model.__dict__['model'] = peft_model + saved = twinkle_model._get_adapter_state_dict_for_save('default') + + assert set(saved) == set(get_peft_model_state_dict(peft_model, adapter_name='default')) + assert 'base_model.model.layers.0.self_attn.q_proj.weight' in saved + + peft_model.peft_config['default'].save_pretrained(tmp_path) + save_file({name: value.contiguous() for name, value in saved.items()}, + str(tmp_path / 'adapter_model.safetensors')) + reloaded_base = TinyDecoder(num_layers=1) + reloaded_base.load_state_dict(base_state) + loaded = PeftModel.from_pretrained(reloaded_base, tmp_path) + + assert torch.allclose(loaded(inputs), expected, atol=1e-5) + + +@pytest.mark.parametrize('adapter_name', ['default', 'zero_shot']) +def test_trainable_parameter_filter_includes_full_modules(adapter_name): + from twinkle.model.transformers.transformers import TransformersModel + + config = build_zero_shot_lora_config( + s_lora=['layers.0.mlp.down_proj'], + s_fft=['layers.0.self_attn.q_proj'], + r=4, + lora_alpha=8, + ) + peft_model = get_peft_model(TinyDecoder(num_layers=1), config, adapter_name=adapter_name) + model = object.__new__(TransformersModel) + model.strategy = type('Strategy', (), {'unwrap_model': lambda _self, inner: inner})() + model.__dict__['model'] = peft_model + + selected = model._get_trainable_parameters(adapter_name) + expected = {name for name, param in peft_model.named_parameters() if param.requires_grad} + + assert set(selected) == expected + assert any('.modules_to_save.' in name for name in selected) From 7a7ea0184e43ee1afd69562ccbfecca48d3699d0 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Wed, 12 Aug 2026 11:33:26 +0800 Subject: [PATCH 2/9] refactor: rename zero-shot LoRA to Spectral Hybrid LoRA --- ...o_shot_lora.py => spectral_hybrid_lora.py} | 82 +++++++++---------- ...o_shot_lora.sh => spectral_hybrid_lora.sh} | 18 ++-- ...=> compute_spectral_hybrid_lora_config.py} | 50 +++++------ ...o_shot_lora.py => spectral_hybrid_lora.py} | 60 +++++++------- ...t_lora.py => test_spectral_hybrid_lora.py} | 44 +++++----- 5 files changed, 127 insertions(+), 127 deletions(-) rename cookbook/transformers/{zero_shot_lora.py => spectral_hybrid_lora.py} (65%) rename cookbook/transformers/{zero_shot_lora.sh => spectral_hybrid_lora.sh} (63%) rename scripts/{compute_zero_shot_hybrid_config.py => compute_spectral_hybrid_lora_config.py} (59%) rename src/twinkle/model/transformers/{zero_shot_lora.py => spectral_hybrid_lora.py} (79%) rename tests/transformers/{test_zero_shot_lora.py => test_spectral_hybrid_lora.py} (87%) diff --git a/cookbook/transformers/zero_shot_lora.py b/cookbook/transformers/spectral_hybrid_lora.py similarity index 65% rename from cookbook/transformers/zero_shot_lora.py rename to cookbook/transformers/spectral_hybrid_lora.py index 710b56c17..1dd4d828e 100644 --- a/cookbook/transformers/zero_shot_lora.py +++ b/cookbook/transformers/spectral_hybrid_lora.py @@ -9,12 +9,12 @@ from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.model import TransformersModel -from twinkle.model.transformers.zero_shot_lora import ( - allocate_zero_shot_modules, - build_zero_shot_lora_config, - build_zero_shot_param_groups, - compute_zero_shot_scores, - select_zero_shot_targets, +from twinkle.model.transformers.spectral_hybrid_lora import ( + allocate_spectral_modules, + build_spectral_lora_config, + build_spectral_param_groups, + compute_spectral_scores, + select_spectral_targets, ) logger = get_logger() @@ -37,17 +37,17 @@ def build_dataset(data_slice) -> Dataset: return dataset -def load_zero_shot_config(config_value: str) -> LoraConfig: - """Load an allocation produced by scripts/compute_zero_shot_hybrid_config.py.""" +def load_spectral_config(config_value: str) -> LoraConfig: + """Load an allocation produced by scripts/compute_spectral_hybrid_lora_config.py.""" config_path = Path(config_value).expanduser() if not config_path.is_file(): - raise FileNotFoundError(f'Zero-shot config JSON file not found: {config_path}') + raise FileNotFoundError(f'Spectral config JSON file not found: {config_path}') with config_path.open(encoding='utf-8') as handle: raw_config = json.load(handle) if not isinstance(raw_config, dict): - raise ValueError('Zero-shot config JSON must contain an object.') - if raw_config.get('method') not in (None, 'zero_shot_spectral'): - raise ValueError(f'Unsupported zero-shot config method: {raw_config.get("method")!r}.') + raise ValueError('Spectral config JSON must contain an object.') + if raw_config.get('method') not in (None, 'spectral_hybrid_lora'): + raise ValueError(f'Unsupported spectral config method: {raw_config.get("method")!r}.') def module_list(primary_key, peft_key): value = raw_config.get(primary_key, raw_config.get(peft_key)) @@ -56,18 +56,18 @@ def module_list(primary_key, peft_key): if isinstance(value, str): value = [value] if not isinstance(value, (list, tuple, set)) or not all(isinstance(item, str) for item in value): - raise ValueError(f'Zero-shot config {primary_key} must be a list of module names.') + raise ValueError(f'Spectral config {primary_key} must be a list of module names.') return sorted(set(value)) s_fft = module_list('s_fft', 'modules_to_save') s_lora = module_list('s_lora', 'target_modules') overlap = set(s_fft) & set(s_lora) if overlap: - raise ValueError(f'Zero-shot modules cannot be both FFT and LoRA: {sorted(overlap)}') - r = int(raw_config.get('r', args.extra.get('zero_shot_r', args.lora.lora_r))) - lora_alpha = int(raw_config.get('lora_alpha', args.extra.get('zero_shot_alpha', r * 2))) + raise ValueError(f'Spectral modules cannot be both FFT and LoRA: {sorted(overlap)}') + r = int(raw_config.get('r', args.extra.get('spectral_r', args.lora.lora_r))) + lora_alpha = int(raw_config.get('lora_alpha', args.extra.get('spectral_alpha', r * 2))) lora_dropout = float(raw_config.get('lora_dropout', 0.0)) - return build_zero_shot_lora_config( + return build_spectral_lora_config( s_lora, s_fft, r=r, @@ -78,22 +78,22 @@ def module_list(primary_key, peft_key): def compute_allocation() -> LoraConfig: """Allocate full fine-tuning and LoRA modules from pretrained spectra.""" - r = int(args.extra.get('zero_shot_r', args.lora.lora_r)) - lora_alpha = int(args.extra.get('zero_shot_alpha', r * 2)) - fft_ratio = float(args.extra.get('zero_shot_fft_ratio', 0.1)) - epsilon = float(args.extra.get('zero_shot_epsilon', 1e-12)) + r = int(args.extra.get('spectral_r', args.lora.lora_r)) + lora_alpha = int(args.extra.get('spectral_alpha', r * 2)) + fft_ratio = float(args.extra.get('spectral_fft_ratio', 0.1)) + epsilon = float(args.extra.get('spectral_epsilon', 1e-12)) - logger.info(f'Zero-shot allocation: loading pretrained model ' + logger.info(f'Spectral Hybrid LoRA allocation: loading pretrained model ' f'(r={r}, alpha={lora_alpha}, FFT budget={fft_ratio:.1%})') base_model = TransformersModel(model_id=args.model.model_id) if base_model._memory_efficient_init: - raise ValueError('Zero-shot spectral scoring requires materialized weights; ' + raise ValueError('Spectral scoring requires materialized weights; ' 'disable memory_efficient_init.') target_config = LoraConfig(**args.get_lora_args()) - targets = select_zero_shot_targets(base_model.model, target_config) + targets = select_spectral_targets(base_model.model, target_config) param_counts = {name: module.weight.numel() for name, module in targets.items()} - cache_dir = Path(args.training.output_dir) / 'zero-shot-spectrum-cache' - scores = compute_zero_shot_scores( + cache_dir = Path(args.training.output_dir) / 'spectral-spectrum-cache' + scores = compute_spectral_scores( base_model.model, target_config, r=r, @@ -102,9 +102,9 @@ def compute_allocation() -> LoraConfig: epsilon=epsilon, log_interval=args.training.log_interval, ) - s_fft, s_lora = allocate_zero_shot_modules(scores, param_counts, fft_ratio=fft_ratio) + s_fft, s_lora = allocate_spectral_modules(scores, param_counts, fft_ratio=fft_ratio) ranked = sorted(scores, key=lambda name: (-scores[name], name)) - logger.info('Zero-shot top-10 modules: ' + ', '.join( + logger.info('Spectral Hybrid LoRA top-10 modules: ' + ', '.join( f'{name}=score:{scores[name]:.4f},effective_rank:{scores.metrics[name]["effective_rank"]:.1f},' f'coverage:{scores.metrics[name]["rank_coverage"]:.4f},' f'condition:{scores.metrics[name]["condition_number"]:.2e},' @@ -112,36 +112,36 @@ def compute_allocation() -> LoraConfig: for name in ranked[:10])) fft_params = sum(param_counts[name] for name in s_fft) total_params = sum(param_counts.values()) - logger.info(f'Zero-shot allocation: {len(s_fft)} FFT modules, {len(s_lora)} LoRA modules ' + logger.info(f'Spectral Hybrid LoRA allocation: {len(s_fft)} FFT modules, {len(s_lora)} LoRA modules ' f'({fft_params / total_params:.1%} of candidate params to FFT; cache={cache_dir})') - logger.info(f'Zero-shot FFT modules: {", ".join(s_fft) if s_fft else "(none)"}') + logger.info(f'Spectral FFT modules: {", ".join(s_fft) if s_fft else "(none)"}') del base_model - return build_zero_shot_lora_config(s_lora, s_fft, r=r, lora_alpha=lora_alpha) + return build_spectral_lora_config(s_lora, s_fft, r=r, lora_alpha=lora_alpha) def train() -> None: train_samples = args.training.train_samples or 1000 - config_value = args.extra.get('zero_shot_config') + config_value = args.extra.get('spectral_config') if config_value: - zero_shot_config = load_zero_shot_config(config_value) - logger.info(f'Using supplied zero-shot config; skipping spectral scoring ' - f'({len(zero_shot_config.modules_to_save or [])} FFT modules, ' - f'{len(zero_shot_config.target_modules or [])} LoRA modules)') + spectral_config = load_spectral_config(config_value) + logger.info(f'Using supplied spectral config; skipping spectral scoring ' + f'({len(spectral_config.modules_to_save or [])} FFT modules, ' + f'{len(spectral_config.target_modules or [])} LoRA modules)') else: - zero_shot_config = compute_allocation() + spectral_config = compute_allocation() dataset = build_dataset(range(train_samples)) dataloader = DataLoader(dataset=dataset, batch_size=args.training.batch_size) model = TransformersModel(model_id=args.model.model_id) model.add_adapter_to_model( args.lora.adapter_name, - zero_shot_config, + spectral_config, gradient_accumulation_steps=args.training.gradient_accumulation_steps, ) - param_groups = build_zero_shot_param_groups( + param_groups = build_spectral_param_groups( model.strategy.unwrap_model(model.model), - lr_lora=float(args.extra.get('zero_shot_lr_lora', args.optimizer.learning_rate)), - lr_fft=float(args.extra.get('zero_shot_lr_fft', 1e-6)), + lr_lora=float(args.extra.get('spectral_lr_lora', args.optimizer.learning_rate)), + lr_fft=float(args.extra.get('spectral_lr_fft', 1e-6)), weight_decay=args.optimizer.weight_decay, adapter_name=args.lora.adapter_name, ) diff --git a/cookbook/transformers/zero_shot_lora.sh b/cookbook/transformers/spectral_hybrid_lora.sh similarity index 63% rename from cookbook/transformers/zero_shot_lora.sh rename to cookbook/transformers/spectral_hybrid_lora.sh index a5b9dfd76..5402a59f7 100644 --- a/cookbook/transformers/zero_shot_lora.sh +++ b/cookbook/transformers/spectral_hybrid_lora.sh @@ -1,9 +1,9 @@ #!/bin/sh # Generate the allocation from pretrained spectra, then train the selected FFT/LoRA modules. -# Reuse an offline allocation with: --zero-shot-config ./output/zero_shot/hybrid_config.json +# Reuse an offline allocation with: --spectral-config ./output/spectral_hybrid_lora/config.json CUDA_VISIBLE_DEVICES=0,1,2,3 \ - torchrun --nproc_per_node=4 zero_shot_lora.py \ + torchrun --nproc_per_node=4 spectral_hybrid_lora.py \ --model-id ms://Qwen/Qwen3.5-9B \ --dataset-id data/financial_sft/processed/finqa_tatqa_train_messages.jsonl \ --template-cls Qwen3_5Template \ @@ -15,16 +15,16 @@ CUDA_VISIBLE_DEVICES=0,1,2,3 \ --weight-decay 0.01 \ --gradient-accumulation-steps 4 \ --log-interval 1 \ - --output-dir ./output/zero_shot_lora \ + --output-dir ./output/spectral_hybrid_lora \ --adapter-name default \ --lora-r 64 \ --scheduler-cls CosineWarmupScheduler \ --num-warmup-steps 10 \ --train-samples 1000 \ - --zero-shot-r 64 \ - --zero-shot-alpha 128 \ - --zero-shot-fft-ratio 0.3 \ - --zero-shot-epsilon 1e-12 \ - --zero-shot-lr-fft 1e-6 \ - --zero-shot-lr-lora 2.5e-5 \ + --spectral-r 64 \ + --spectral-alpha 128 \ + --spectral-fft-ratio 0.3 \ + --spectral-epsilon 1e-12 \ + --spectral-lr-fft 1e-6 \ + --spectral-lr-lora 2.5e-5 \ "$@" diff --git a/scripts/compute_zero_shot_hybrid_config.py b/scripts/compute_spectral_hybrid_lora_config.py similarity index 59% rename from scripts/compute_zero_shot_hybrid_config.py rename to scripts/compute_spectral_hybrid_lora_config.py index 12fd60cdf..03ba4ad6a 100644 --- a/scripts/compute_zero_shot_hybrid_config.py +++ b/scripts/compute_spectral_hybrid_lora_config.py @@ -1,13 +1,13 @@ #!/usr/bin/env python3 -"""Compute a zero-shot FFT/LoRA allocation from pretrained weight spectra. +"""Compute a data-free spectral FFT/LoRA allocation from pretrained spectra. Example: - python scripts/compute_zero_shot_hybrid_config.py \ + python scripts/compute_spectral_hybrid_lora_config.py \ --model-id ms://Qwen/Qwen3.5-9B \ - --zero-shot-r 64 \ - --zero-shot-fft-ratio 0.3 \ - --output-dir ./output/zero_shot \ - --hybrid-config-output ./output/zero_shot/hybrid_config.json + --spectral-r 64 \ + --spectral-fft-ratio 0.3 \ + --output-dir ./output/spectral_hybrid_lora \ + --spectral-config-output ./output/spectral_hybrid_lora/config.json """ import json @@ -19,8 +19,8 @@ from twinkle import DeviceMesh, Platform, get_logger from twinkle.cli import CLI from twinkle.model import TransformersModel -from twinkle.model.transformers.zero_shot_lora import (CANDIDATE_TYPES, allocate_zero_shot_modules, - compute_zero_shot_scores, select_zero_shot_targets) +from twinkle.model.transformers.spectral_hybrid_lora import (CANDIDATE_TYPES, allocate_spectral_modules, + compute_spectral_scores, select_spectral_targets) logger = get_logger() args = CLI.from_args() @@ -34,20 +34,20 @@ def main() -> None: if not args.model.model_id: raise ValueError('--model-id is required.') - r = int(args.extra.get('zero_shot_r', args.lora.lora_r)) - lora_alpha = int(args.extra.get('zero_shot_alpha', r * 2)) - fft_ratio = float(args.extra.get('zero_shot_fft_ratio', 0.1)) - epsilon = float(args.extra.get('zero_shot_epsilon', 1e-12)) - output_path = Path(args.extra.get('hybrid_config_output', - Path(args.training.output_dir) / 'hybrid_config.json')).expanduser() - cache_dir = Path( - args.extra.get('zero_shot_cache_dir', - Path(args.training.output_dir) / 'zero-shot-spectrum-cache')).expanduser() + r = int(args.extra.get('spectral_r', args.lora.lora_r)) + lora_alpha = int(args.extra.get('spectral_alpha', r * 2)) + fft_ratio = float(args.extra.get('spectral_fft_ratio', 0.1)) + epsilon = float(args.extra.get('spectral_epsilon', 1e-12)) + output_path = Path( + args.extra.get('spectral_config_output', + Path(args.training.output_dir) / 'spectral_hybrid_lora_config.json')).expanduser() + cache_dir = Path(args.extra.get('spectral_cache_dir', + Path(args.training.output_dir) / 'spectral-spectrum-cache')).expanduser() - logger.info(f'Loading pretrained model for zero-shot spectral scoring: {args.model.model_id}') + logger.info(f'Loading pretrained model for Spectral Hybrid LoRA scoring: {args.model.model_id}') model = TransformersModel(model_id=args.model.model_id) if model._memory_efficient_init: - raise ValueError('Zero-shot spectral scoring requires materialized weights; disable memory_efficient_init.') + raise ValueError('Spectral scoring requires materialized weights; disable memory_efficient_init.') target_config = LoraConfig( r=r, @@ -55,9 +55,9 @@ def main() -> None: lora_dropout=0.0, target_modules=list(CANDIDATE_TYPES.values()), ) - targets = select_zero_shot_targets(model.model, target_config) + targets = select_spectral_targets(model.model, target_config) param_counts = {name: module.weight.numel() for name, module in targets.items()} - scores = compute_zero_shot_scores( + scores = compute_spectral_scores( model.model, target_config, r=r, @@ -67,12 +67,12 @@ def main() -> None: log_interval=args.training.log_interval, ) counts = {name: param_counts[name] for name in scores} - s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=fft_ratio) + s_fft, s_lora = allocate_spectral_modules(scores, counts, fft_ratio=fft_ratio) fft_params = sum(counts[name] for name in s_fft) total_params = sum(counts.values()) config = { - 'method': 'zero_shot_spectral', + 'method': 'spectral_hybrid_lora', 'model_id': args.model.model_id, 's_fft': s_fft, 's_lora': s_lora, @@ -81,7 +81,7 @@ def main() -> None: 'lora_dropout': 0.0, 'fft_ratio': fft_ratio, 'realized_fft_param_ratio': fft_params / total_params, - 'zero_shot_epsilon': epsilon, + 'spectral_epsilon': epsilon, 'metrics': { name: { 'score': scores[name], @@ -98,7 +98,7 @@ def main() -> None: json.dump(config, handle, ensure_ascii=False, indent=2, sort_keys=True) handle.write('\n') os.replace(temporary_path, output_path) - logger.info(f'Zero-shot config written to {output_path}: ' + logger.info(f'Spectral config written to {output_path}: ' f'{len(s_fft)} FFT modules, {len(s_lora)} LoRA modules, ' f'realized FFT parameter ratio={fft_params / total_params:.2%}') diff --git a/src/twinkle/model/transformers/zero_shot_lora.py b/src/twinkle/model/transformers/spectral_hybrid_lora.py similarity index 79% rename from src/twinkle/model/transformers/zero_shot_lora.py rename to src/twinkle/model/transformers/spectral_hybrid_lora.py index 3c9cf5741..92faf3043 100644 --- a/src/twinkle/model/transformers/zero_shot_lora.py +++ b/src/twinkle/model/transformers/spectral_hybrid_lora.py @@ -28,26 +28,26 @@ _LAYER_RE = re.compile(r'\blayers\.(\d+)\.') -class ZeroShotScores(dict[str, float]): - """Zero-shot spectral scores with per-module metric details.""" +class SpectralScores(dict[str, float]): + """Spectral scores with per-module metric details.""" def __init__(self, scores: Mapping[str, float], metrics: Mapping[str, Mapping[str, float]]) -> None: super().__init__(scores) self.metrics = {name: dict(values) for name, values in metrics.items()} -def compute_zero_shot_spectral_metrics( +def compute_spectral_metrics( singular_values: torch.Tensor, r: int, epsilon: float = 1e-12, ) -> Dict[str, float]: - """Compute the pretrained-spectrum metrics used for zero-shot allocation.""" + """Compute pretrained-spectrum metrics for data-free spectral allocation.""" if r <= 0: - raise ValueError('Zero-shot LoRA rank r must be positive.') + raise ValueError('Spectral LoRA rank r must be positive.') if singular_values.ndim != 1 or singular_values.numel() == 0: - raise ValueError('Zero-shot scoring requires a non-empty singular-value vector.') + raise ValueError('Spectral scoring requires a non-empty singular-value vector.') if epsilon <= 0: - raise ValueError('Zero-shot scoring epsilon must be positive.') + raise ValueError('Spectral scoring epsilon must be positive.') values = singular_values.detach().to(device='cpu', dtype=torch.float64).abs() values = values.sort(descending=True).values @@ -84,8 +84,8 @@ def _is_candidate_module(module_name: str) -> bool: return _LAYER_RE.search(module_name) is not None and module_name.rsplit('.', 1)[-1] in _CANDIDATE_SUFFIXES -def select_zero_shot_targets(model: nn.Module, config: LoraConfig) -> Dict[str, nn.Module]: - """Resolve materialized linear modules eligible for zero-shot allocation.""" +def select_spectral_targets(model: nn.Module, config: LoraConfig) -> Dict[str, nn.Module]: + """Resolve materialized linear modules eligible for spectral allocation.""" resolved = _maybe_include_all_linear_layers(copy.deepcopy(config), model) targets: Dict[str, nn.Module] = {} for name, module in model.named_modules(): @@ -95,13 +95,13 @@ def select_zero_shot_targets(model: nn.Module, config: LoraConfig) -> Dict[str, if not check_target_module_exists(resolved, name) or not _is_candidate_module(name): continue if not isinstance(module, nn.Linear): - raise ValueError(f'Zero-shot LoRA target {name!r} is not an nn.Linear module.') + raise ValueError(f'Spectral LoRA target {name!r} is not an nn.Linear module.') if weight.is_meta: - raise ValueError('Zero-shot spectral scoring requires materialized weights; ' + raise ValueError('Spectral scoring requires materialized weights; ' 'disable memory_efficient_init.') targets[name] = module if not targets: - raise ValueError(f'Zero-shot LoRA found no candidate modules for {config.target_modules!r}.') + raise ValueError(f'Spectral LoRA found no candidate modules for {config.target_modules!r}.') return targets @@ -118,7 +118,7 @@ def _spectrum_cache_path( @torch.no_grad() -def compute_zero_shot_scores( +def compute_spectral_scores( model: nn.Module, config: LoraConfig, r: int, @@ -127,9 +127,9 @@ def compute_zero_shot_scores( cache_key: str = '', epsilon: float = 1e-12, log_interval: int = 20, -) -> ZeroShotScores: - """Score pretrained modules from singular-value spectra without training data.""" - targets = select_zero_shot_targets(model, config) +) -> SpectralScores: + """Score pretrained modules from singular-value spectra for data-free allocation.""" + targets = select_spectral_targets(model, config) names = sorted(targets) distributed = dist.is_available() and dist.is_initialized() rank = dist.get_rank() if distributed else 0 @@ -138,7 +138,7 @@ def compute_zero_shot_scores( if rank == 0: if log_interval > 0: - logger.info(f'Zero-shot spectral scoring: {len(names)} modules, LoRA rank={r}') + logger.info(f'Spectral Hybrid LoRA scoring: {len(names)} modules, LoRA rank={r}') for index, name in enumerate(names, start=1): weight = targets[name].weight.detach() shape = tuple(weight.shape) @@ -153,7 +153,7 @@ def compute_zero_shot_scores( if isinstance(cached, torch.Tensor) and cached.ndim == 1 and cached.numel() == min(shape): singular_values = cached except (OSError, RuntimeError, EOFError): - logger.warning(f'Zero-shot spectrum cache is unreadable: {spectrum_path}; recomputing') + logger.warning(f'Spectral spectrum cache is unreadable: {spectrum_path}; recomputing') if singular_values is None: cpu_weight = weight.to(device='cpu') @@ -166,11 +166,11 @@ def compute_zero_shot_scores( torch.save(singular_values, temporary_path) os.replace(temporary_path, spectrum_path) - metrics = compute_zero_shot_spectral_metrics(singular_values, r=r, epsilon=epsilon) + metrics = compute_spectral_metrics(singular_values, r=r, epsilon=epsilon) metrics_by_module[name] = metrics scores[name] = metrics['score'] if log_interval > 0 and (index % log_interval == 0 or index == len(names)): - logger.info(f'Zero-shot spectral scoring: {index}/{len(names)} modules ' + logger.info(f'Spectral Hybrid LoRA scoring: {index}/{len(names)} modules ' f'(last {name} -> score={metrics["score"]:.4f}, ' f'effective_rank={metrics["effective_rank"]:.1f}, ' f'rank_coverage={metrics["rank_coverage"]:.4f})') @@ -179,23 +179,23 @@ def compute_zero_shot_scores( payload = [(scores, metrics_by_module) if rank == 0 else None] dist.broadcast_object_list(payload, src=0) scores, metrics_by_module = payload[0] - return ZeroShotScores(scores, metrics_by_module) + return SpectralScores(scores, metrics_by_module) -def allocate_zero_shot_modules( +def allocate_spectral_modules( scores: Mapping[str, float], param_counts: Mapping[str, int], fft_ratio: float = 0.1, ) -> Tuple[List[str], List[str]]: """Allocate the highest-scoring prefix to full fine-tuning within a parameter budget.""" if not 0.0 <= fft_ratio < 1.0: - raise ValueError('Zero-shot FFT ratio must be in [0, 1).') + raise ValueError('Spectral FFT ratio must be in [0, 1).') missing = set(scores) - set(param_counts) if missing: - raise ValueError(f'Zero-shot allocation is missing parameter counts for: {sorted(missing)}.') + raise ValueError(f'Spectral allocation is missing parameter counts for: {sorted(missing)}.') total = sum(param_counts[name] for name in scores) if total <= 0: - raise ValueError('Zero-shot candidate parameter count must be positive.') + raise ValueError('Spectral candidate parameter count must be positive.') budget = fft_ratio * total ordered = sorted(scores, key=lambda name: (-scores[name], name)) @@ -212,7 +212,7 @@ def allocate_zero_shot_modules( return sorted(s_fft), sorted(s_lora) -def build_zero_shot_lora_config( +def build_spectral_lora_config( s_lora: List[str], s_fft: List[str], r: int = 16, @@ -221,7 +221,7 @@ def build_zero_shot_lora_config( **kwargs, ) -> LoraConfig: if not s_lora: - raise ValueError('Zero-shot LoRA requires at least one LoRA module.') + raise ValueError('Spectral LoRA requires at least one LoRA module.') return LoraConfig( r=r, lora_alpha=lora_alpha, @@ -232,7 +232,7 @@ def build_zero_shot_lora_config( ) -def build_zero_shot_param_groups( +def build_spectral_param_groups( peft_model: nn.Module, lr_lora: float = 2.5e-5, lr_fft: float = 1e-6, @@ -253,7 +253,7 @@ def build_zero_shot_param_groups( fft_params.append(param) fft_names.append(name) else: - raise ValueError(f'Zero-shot LoRA cannot classify trainable parameter {name!r}.') + raise ValueError(f'Spectral LoRA cannot classify trainable parameter {name!r}.') groups: List[dict] = [] if lora_params: @@ -261,5 +261,5 @@ def build_zero_shot_param_groups( if fft_params: groups.append({'params': fft_params, 'param_names': fft_names, 'lr': lr_fft, 'weight_decay': weight_decay}) if not groups: - raise ValueError('Zero-shot LoRA found no trainable parameters to optimize.') + raise ValueError('Spectral LoRA found no trainable parameters to optimize.') return groups diff --git a/tests/transformers/test_zero_shot_lora.py b/tests/transformers/test_spectral_hybrid_lora.py similarity index 87% rename from tests/transformers/test_zero_shot_lora.py rename to tests/transformers/test_spectral_hybrid_lora.py index 7db4307d9..27f25cdd2 100644 --- a/tests/transformers/test_zero_shot_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -4,14 +4,14 @@ from peft.utils import get_peft_model_state_dict from torch import nn -from twinkle.model.transformers.zero_shot_lora import ( +from twinkle.model.transformers.spectral_hybrid_lora import ( CANDIDATE_TYPES, - allocate_zero_shot_modules, - build_zero_shot_lora_config, - build_zero_shot_param_groups, - compute_zero_shot_scores, - compute_zero_shot_spectral_metrics, - select_zero_shot_targets, + allocate_spectral_modules, + build_spectral_lora_config, + build_spectral_param_groups, + compute_spectral_scores, + compute_spectral_metrics, + select_spectral_targets, ) @@ -42,7 +42,7 @@ def forward(self, inputs): def test_spectral_metrics_match_weighted_formula(): singular_values = torch.tensor([4.0, 2.0, 1.0, 0.5], dtype=torch.float64) - metrics = compute_zero_shot_spectral_metrics(singular_values, r=1) + metrics = compute_spectral_metrics(singular_values, r=1) probabilities = singular_values / singular_values.sum() effective_rank = torch.exp(-(probabilities * probabilities.log()).sum()).item() @@ -69,14 +69,14 @@ def test_spectral_metrics_match_weighted_formula(): ]) def test_spectral_metrics_validate_inputs(singular_values, r, message): with pytest.raises(ValueError, match=message): - compute_zero_shot_spectral_metrics(singular_values, r=r) + compute_spectral_metrics(singular_values, r=r) def test_select_targets_covers_supported_module_types(): model = TinyDecoder(num_layers=2) config = LoraConfig(r=4, target_modules=list(CANDIDATE_TYPES.values())) - targets = select_zero_shot_targets(model, config) + targets = select_spectral_targets(model, config) assert len(targets) == 14 assert all(name.rsplit('.', 1)[-1] in CANDIDATE_TYPES.values() for name in targets) @@ -93,9 +93,9 @@ def record_svdvals(weight): return original_svdvals(weight) monkeypatch.setattr(torch.linalg, 'svdvals', record_svdvals) - scores = compute_zero_shot_scores( + scores = compute_spectral_scores( model, config, r=2, cache_dir=tmp_path, cache_key='tiny', log_interval=0) - cached_scores = compute_zero_shot_scores( + cached_scores = compute_spectral_scores( model, config, r=2, cache_dir=tmp_path, cache_key='tiny', log_interval=0) assert set(scores) == {'layers.0.mlp.down_proj', 'layers.0.self_attn.q_proj'} @@ -109,7 +109,7 @@ def test_allocation_prioritizes_high_scores_within_budget(): scores = {'a': 0.2, 'b': 0.8, 'c': 0.4, 'd': 0.6} counts = {name: 100 for name in scores} - s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=0.25) + s_fft, s_lora = allocate_spectral_modules(scores, counts, fft_ratio=0.25) assert s_fft == ['b'] assert set(s_lora) == {'a', 'c', 'd'} @@ -119,7 +119,7 @@ def test_allocation_uses_strict_ranked_prefix(): scores = {'big': 0.9, 'small_a': 0.8, 'small_b': 0.7} counts = {'big': 80, 'small_a': 30, 'small_b': 15} - s_fft, s_lora = allocate_zero_shot_modules(scores, counts, fft_ratio=0.7) + s_fft, s_lora = allocate_spectral_modules(scores, counts, fft_ratio=0.7) assert s_fft == ['big'] assert set(s_lora) == {'small_a', 'small_b'} @@ -127,16 +127,16 @@ def test_allocation_uses_strict_ranked_prefix(): def test_allocation_rejects_full_fft_budget(): with pytest.raises(ValueError, match=r'\[0, 1\)'): - allocate_zero_shot_modules({'module': 1.0}, {'module': 10}, fft_ratio=1.0) + allocate_spectral_modules({'module': 1.0}, {'module': 10}, fft_ratio=1.0) def test_config_requires_a_lora_target(): with pytest.raises(ValueError, match='at least one LoRA module'): - build_zero_shot_lora_config([], ['layers.0.self_attn.q_proj']) + build_spectral_lora_config([], ['layers.0.self_attn.q_proj']) def test_config_and_param_groups_cover_every_trainable_parameter(): - config = build_zero_shot_lora_config( + config = build_spectral_lora_config( s_lora=['layers.0.mlp.down_proj'], s_fft=['layers.0.self_attn.q_proj'], r=4, @@ -144,7 +144,7 @@ def test_config_and_param_groups_cover_every_trainable_parameter(): ) model = get_peft_model(TinyDecoder(num_layers=1), config) - groups = build_zero_shot_param_groups(model, lr_lora=2.5e-5, lr_fft=1e-6) + groups = build_spectral_param_groups(model, lr_lora=2.5e-5, lr_fft=1e-6) assert {group['lr'] for group in groups} == {2.5e-5, 1e-6} grouped = {id(param) for group in groups for param in group['params']} @@ -166,7 +166,7 @@ def test_strategy_adapter_state_includes_full_modules(strategy_cls): strategy = object.__new__(Strategy) strategy.ep_fsdp_device_mesh = None - config = build_zero_shot_lora_config( + config = build_spectral_lora_config( s_lora=['layers.0.mlp.down_proj'], s_fft=['layers.0.self_attn.q_proj'], r=4, @@ -188,7 +188,7 @@ def test_twinkle_checkpoint_normalization_round_trips_full_modules(tmp_path): torch.manual_seed(0) base = TinyDecoder(num_layers=1) base_state = {name: value.detach().clone() for name, value in base.state_dict().items()} - config = build_zero_shot_lora_config( + config = build_spectral_lora_config( s_lora=['layers.0.mlp.down_proj'], s_fft=['layers.0.self_attn.q_proj'], r=4, @@ -222,11 +222,11 @@ def test_twinkle_checkpoint_normalization_round_trips_full_modules(tmp_path): assert torch.allclose(loaded(inputs), expected, atol=1e-5) -@pytest.mark.parametrize('adapter_name', ['default', 'zero_shot']) +@pytest.mark.parametrize('adapter_name', ['default', 'spectral_hybrid']) def test_trainable_parameter_filter_includes_full_modules(adapter_name): from twinkle.model.transformers.transformers import TransformersModel - config = build_zero_shot_lora_config( + config = build_spectral_lora_config( s_lora=['layers.0.mlp.down_proj'], s_fft=['layers.0.self_attn.q_proj'], r=4, From 8a91cb10eb53d90722f77adbf49bb11d02959f4c Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Wed, 12 Aug 2026 12:32:50 +0800 Subject: [PATCH 3/9] refactor: integrate spectral allocation into cookbook --- cookbook/transformers/spectral_hybrid_lora.py | 100 +++++++++++----- cookbook/transformers/spectral_hybrid_lora.sh | 4 +- .../compute_spectral_hybrid_lora_config.py | 107 ------------------ src/twinkle/model/base.py | 62 +++++----- .../transformers/spectral_hybrid_lora.py | 22 +++- .../transformers/test_spectral_hybrid_lora.py | 67 +++++++++++ 6 files changed, 193 insertions(+), 169 deletions(-) delete mode 100644 scripts/compute_spectral_hybrid_lora_config.py diff --git a/cookbook/transformers/spectral_hybrid_lora.py b/cookbook/transformers/spectral_hybrid_lora.py index 1dd4d828e..87e79330f 100644 --- a/cookbook/transformers/spectral_hybrid_lora.py +++ b/cookbook/transformers/spectral_hybrid_lora.py @@ -1,19 +1,24 @@ import json +import os from pathlib import Path +import torch.distributed as dist from peft import LoraConfig import twinkle -from twinkle import DeviceMesh, get_device_placement, get_logger +from twinkle import DeviceMesh, Platform, get_device_placement, get_logger from twinkle.cli import CLI from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.model import TransformersModel +from twinkle.model.base import initialize_process_group from twinkle.model.transformers.spectral_hybrid_lora import ( + CANDIDATE_TYPES, allocate_spectral_modules, build_spectral_lora_config, build_spectral_param_groups, compute_spectral_scores, + resolve_spectral_config_path, select_spectral_targets, ) @@ -37,11 +42,8 @@ def build_dataset(data_slice) -> Dataset: return dataset -def load_spectral_config(config_value: str) -> LoraConfig: - """Load an allocation produced by scripts/compute_spectral_hybrid_lora_config.py.""" - config_path = Path(config_value).expanduser() - if not config_path.is_file(): - raise FileNotFoundError(f'Spectral config JSON file not found: {config_path}') +def load_spectral_config(config_path: Path) -> LoraConfig: + """Load an existing Spectral Hybrid LoRA allocation.""" with config_path.open(encoding='utf-8') as handle: raw_config = json.load(handle) if not isinstance(raw_config, dict): @@ -76,8 +78,8 @@ def module_list(primary_key, peft_key): ) -def compute_allocation() -> LoraConfig: - """Allocate full fine-tuning and LoRA modules from pretrained spectra.""" +def compute_allocation(config_path: Path) -> None: + """Compute and persist a data-free spectral allocation on the master rank.""" r = int(args.extra.get('spectral_r', args.lora.lora_r)) lora_alpha = int(args.extra.get('spectral_alpha', r * 2)) fft_ratio = float(args.extra.get('spectral_fft_ratio', 0.1)) @@ -85,14 +87,22 @@ def compute_allocation() -> LoraConfig: logger.info(f'Spectral Hybrid LoRA allocation: loading pretrained model ' f'(r={r}, alpha={lora_alpha}, FFT budget={fft_ratio:.1%})') - base_model = TransformersModel(model_id=args.model.model_id) + analysis_mesh = DeviceMesh.from_sizes(world_size=1, dp_size=1) + base_model = TransformersModel(model_id=args.model.model_id, device_mesh=analysis_mesh) if base_model._memory_efficient_init: raise ValueError('Spectral scoring requires materialized weights; ' 'disable memory_efficient_init.') - target_config = LoraConfig(**args.get_lora_args()) + target_modules = args.lora.lora_target_modules or list(CANDIDATE_TYPES.values()) + target_config = LoraConfig( + r=r, + lora_alpha=lora_alpha, + lora_dropout=0.0, + target_modules=target_modules, + ) targets = select_spectral_targets(base_model.model, target_config) param_counts = {name: module.weight.numel() for name, module in targets.items()} - cache_dir = Path(args.training.output_dir) / 'spectral-spectrum-cache' + cache_dir = Path(args.extra.get( + 'spectral_cache_dir', Path(args.training.output_dir) / 'spectral-spectrum-cache')).expanduser() scores = compute_spectral_scores( base_model.model, target_config, @@ -101,34 +111,72 @@ def compute_allocation() -> LoraConfig: cache_key=str(args.model.model_id), epsilon=epsilon, log_interval=args.training.log_interval, + broadcast=False, ) s_fft, s_lora = allocate_spectral_modules(scores, param_counts, fft_ratio=fft_ratio) - ranked = sorted(scores, key=lambda name: (-scores[name], name)) - logger.info('Spectral Hybrid LoRA top-10 modules: ' + ', '.join( - f'{name}=score:{scores[name]:.4f},effective_rank:{scores.metrics[name]["effective_rank"]:.1f},' - f'coverage:{scores.metrics[name]["rank_coverage"]:.4f},' - f'condition:{scores.metrics[name]["condition_number"]:.2e},' - f'decay:{scores.metrics[name]["decay"]:.4e}' - for name in ranked[:10])) fft_params = sum(param_counts[name] for name in s_fft) total_params = sum(param_counts.values()) + realized_fft_param_ratio = fft_params / total_params logger.info(f'Spectral Hybrid LoRA allocation: {len(s_fft)} FFT modules, {len(s_lora)} LoRA modules ' - f'({fft_params / total_params:.1%} of candidate params to FFT; cache={cache_dir})') + f'({realized_fft_param_ratio:.1%} of candidate params to FFT; cache={cache_dir})') logger.info(f'Spectral FFT modules: {", ".join(s_fft) if s_fft else "(none)"}') + + raw_config = { + 'method': 'spectral_hybrid_lora', + 'model_id': args.model.model_id, + 's_fft': s_fft, + 's_lora': s_lora, + 'r': r, + 'lora_alpha': lora_alpha, + 'lora_dropout': 0.0, + 'fft_ratio': fft_ratio, + 'realized_fft_param_ratio': realized_fft_param_ratio, + 'spectral_epsilon': epsilon, + 'metrics': { + name: { + **scores.metrics[name], + 'score': scores[name], + } + for name in sorted(scores) + }, + } + if Platform.is_master(): + config_path.parent.mkdir(parents=True, exist_ok=True) + temporary_path = config_path.with_suffix(f'{config_path.suffix}.tmp') + with temporary_path.open('w', encoding='utf-8') as handle: + json.dump(raw_config, handle, ensure_ascii=False, indent=2, sort_keys=True) + handle.write('\n') + os.replace(temporary_path, config_path) + logger.info(f'Spectral config written to {config_path}') + del base_model - return build_spectral_lora_config(s_lora, s_fft, r=r, lora_alpha=lora_alpha) -def train() -> None: - train_samples = args.training.train_samples or 1000 +def resolve_spectral_config() -> LoraConfig: config_value = args.extra.get('spectral_config') - if config_value: - spectral_config = load_spectral_config(config_value) - logger.info(f'Using supplied spectral config; skipping spectral scoring ' + config_path, should_load = resolve_spectral_config_path(config_value, Path(args.training.output_dir)) + if should_load: + spectral_config = load_spectral_config(config_path) + logger.info(f'Using existing spectral config {config_path}; skipping spectral scoring ' f'({len(spectral_config.modules_to_save or [])} FFT modules, ' f'{len(spectral_config.target_modules or [])} LoRA modules)') + return spectral_config + if config_value: + logger.info(f'Spectral config {config_path} does not exist; computing it from model weights') else: - spectral_config = compute_allocation() + logger.info(f'No spectral config supplied; computing allocation and writing {config_path}') + + initialize_process_group() + if Platform.is_master(): + compute_allocation(config_path) + if dist.is_available() and dist.is_initialized(): + dist.barrier() + return load_spectral_config(config_path) + + +def train() -> None: + train_samples = args.training.train_samples or 1000 + spectral_config = resolve_spectral_config() dataset = build_dataset(range(train_samples)) dataloader = DataLoader(dataset=dataset, batch_size=args.training.batch_size) diff --git a/cookbook/transformers/spectral_hybrid_lora.sh b/cookbook/transformers/spectral_hybrid_lora.sh index 5402a59f7..e115ab574 100644 --- a/cookbook/transformers/spectral_hybrid_lora.sh +++ b/cookbook/transformers/spectral_hybrid_lora.sh @@ -1,6 +1,5 @@ #!/bin/sh -# Generate the allocation from pretrained spectra, then train the selected FFT/LoRA modules. -# Reuse an offline allocation with: --spectral-config ./output/spectral_hybrid_lora/config.json +# Reuse the configured allocation, or compute and persist it when the file does not exist. CUDA_VISIBLE_DEVICES=0,1,2,3 \ torchrun --nproc_per_node=4 spectral_hybrid_lora.py \ @@ -16,6 +15,7 @@ CUDA_VISIBLE_DEVICES=0,1,2,3 \ --gradient-accumulation-steps 4 \ --log-interval 1 \ --output-dir ./output/spectral_hybrid_lora \ + --spectral-config ./output/spectral_hybrid_lora/config.json \ --adapter-name default \ --lora-r 64 \ --scheduler-cls CosineWarmupScheduler \ diff --git a/scripts/compute_spectral_hybrid_lora_config.py b/scripts/compute_spectral_hybrid_lora_config.py deleted file mode 100644 index 03ba4ad6a..000000000 --- a/scripts/compute_spectral_hybrid_lora_config.py +++ /dev/null @@ -1,107 +0,0 @@ -#!/usr/bin/env python3 -"""Compute a data-free spectral FFT/LoRA allocation from pretrained spectra. - -Example: - python scripts/compute_spectral_hybrid_lora_config.py \ - --model-id ms://Qwen/Qwen3.5-9B \ - --spectral-r 64 \ - --spectral-fft-ratio 0.3 \ - --output-dir ./output/spectral_hybrid_lora \ - --spectral-config-output ./output/spectral_hybrid_lora/config.json -""" - -import json -import os -from pathlib import Path -from peft import LoraConfig - -import twinkle -from twinkle import DeviceMesh, Platform, get_logger -from twinkle.cli import CLI -from twinkle.model import TransformersModel -from twinkle.model.transformers.spectral_hybrid_lora import (CANDIDATE_TYPES, allocate_spectral_modules, - compute_spectral_scores, select_spectral_targets) - -logger = get_logger() -args = CLI.from_args() - -# This utility is intentionally single-process: it only reads weights and writes one JSON config. -device_mesh = DeviceMesh.from_sizes(fsdp_size=1, dp_size=1) -twinkle.initialize(mode=args.infra.mode, global_device_mesh=device_mesh) - - -def main() -> None: - if not args.model.model_id: - raise ValueError('--model-id is required.') - - r = int(args.extra.get('spectral_r', args.lora.lora_r)) - lora_alpha = int(args.extra.get('spectral_alpha', r * 2)) - fft_ratio = float(args.extra.get('spectral_fft_ratio', 0.1)) - epsilon = float(args.extra.get('spectral_epsilon', 1e-12)) - output_path = Path( - args.extra.get('spectral_config_output', - Path(args.training.output_dir) / 'spectral_hybrid_lora_config.json')).expanduser() - cache_dir = Path(args.extra.get('spectral_cache_dir', - Path(args.training.output_dir) / 'spectral-spectrum-cache')).expanduser() - - logger.info(f'Loading pretrained model for Spectral Hybrid LoRA scoring: {args.model.model_id}') - model = TransformersModel(model_id=args.model.model_id) - if model._memory_efficient_init: - raise ValueError('Spectral scoring requires materialized weights; disable memory_efficient_init.') - - target_config = LoraConfig( - r=r, - lora_alpha=lora_alpha, - lora_dropout=0.0, - target_modules=list(CANDIDATE_TYPES.values()), - ) - targets = select_spectral_targets(model.model, target_config) - param_counts = {name: module.weight.numel() for name, module in targets.items()} - scores = compute_spectral_scores( - model.model, - target_config, - r=r, - cache_dir=cache_dir, - cache_key=str(args.model.model_id), - epsilon=epsilon, - log_interval=args.training.log_interval, - ) - counts = {name: param_counts[name] for name in scores} - s_fft, s_lora = allocate_spectral_modules(scores, counts, fft_ratio=fft_ratio) - fft_params = sum(counts[name] for name in s_fft) - total_params = sum(counts.values()) - - config = { - 'method': 'spectral_hybrid_lora', - 'model_id': args.model.model_id, - 's_fft': s_fft, - 's_lora': s_lora, - 'r': r, - 'lora_alpha': lora_alpha, - 'lora_dropout': 0.0, - 'fft_ratio': fft_ratio, - 'realized_fft_param_ratio': fft_params / total_params, - 'spectral_epsilon': epsilon, - 'metrics': { - name: { - 'score': scores[name], - **scores.metrics[name], - } - for name in sorted(scores) - }, - } - - if Platform.is_master(): - output_path.parent.mkdir(parents=True, exist_ok=True) - temporary_path = output_path.with_suffix(f'{output_path.suffix}.tmp') - with temporary_path.open('w', encoding='utf-8') as handle: - json.dump(config, handle, ensure_ascii=False, indent=2, sort_keys=True) - handle.write('\n') - os.replace(temporary_path, output_path) - logger.info(f'Spectral config written to {output_path}: ' - f'{len(s_fft)} FFT modules, {len(s_lora)} LoRA modules, ' - f'realized FFT parameter ratio={fft_params / total_params:.2%}') - - -if __name__ == '__main__': - main() diff --git a/src/twinkle/model/base.py b/src/twinkle/model/base.py index 8ea00d696..d9e7dcbf1 100644 --- a/src/twinkle/model/base.py +++ b/src/twinkle/model/base.py @@ -18,6 +18,37 @@ from torch.optim.lr_scheduler import LRScheduler +def initialize_process_group(should_bind_device_id: Optional[Callable[[str], bool]] = None) -> None: + """Initialize Twinkle's default distributed process group when launched with multiple ranks.""" + import torch + import torch.distributed as dist + if dist.is_initialized() or Platform.get_world_size() <= 1: + return + + torch_util.set_device() + backend = Platform.device_backend() + if backend == 'hccl': + # Keep training-side HCCL sockets on a per-job port layout to avoid collisions. + from twinkle.utils.platforms import ensure_hccl_socket_env + master_port = int(os.environ.get('MASTER_PORT', '29500')) + ensure_hccl_socket_env(master_port) + init_kwargs = { + 'backend': backend, + 'init_method': 'env://', + 'rank': Platform.get_rank(), + 'world_size': Platform.get_world_size(), + } + bind_device = should_bind_device_id(backend) if should_bind_device_id else backend in ('nccl', 'hccl') + if bind_device: + init_kwargs['device_id'] = torch.device(Platform.get_local_device()) + dist.init_process_group(**init_kwargs) + if backend == 'hccl': + # A bound HCCL default group can leak its device into later Gloo metric groups. + default_pg = dist.distributed_c10d._get_default_group() + if getattr(default_pg, 'bound_device_id', None) is not None: + default_pg.bound_device_id = None + + class TwinkleModel(ABC): _checkpoint_engine = None @@ -146,33 +177,4 @@ def _should_bind_device_id_for_process_group(self, backend: str) -> bool: return backend in ('nccl', 'hccl') def _try_init_process_group(self): - import torch - import torch.distributed as dist - if not dist.is_initialized() and Platform.get_world_size() > 1: - torch_util.set_device() - backend = Platform.device_backend() - if backend == 'hccl': - # fix: In multi-job NPU runs, HCCL default ports may collide (bind/listen failures). - # fix: Inject deterministic per-job port ranges before PG init to reduce cross-job conflicts. - # Keep training-side HCCL sockets on a per-job port layout to - # avoid collisions with other jobs on the same host. - from twinkle.utils.platforms import ensure_hccl_socket_env - master_port = int(os.environ.get('MASTER_PORT', '29500')) - ensure_hccl_socket_env(master_port) - init_kwargs = { - 'backend': backend, - 'init_method': 'env://', - 'rank': Platform.get_rank(), - 'world_size': Platform.get_world_size(), - } - if self._should_bind_device_id_for_process_group(backend): - init_kwargs['device_id'] = torch.device(Platform.get_local_device()) - dist.init_process_group(**init_kwargs) - if backend == 'hccl': - default_pg = dist.distributed_c10d._get_default_group() - if getattr(default_pg, 'bound_device_id', None) is not None: - # If the default HCCL PG keeps a bound device id, PyTorch may - # propagate that binding into later Gloo subgroup creation. That - # breaks the metrics/object-gather path on NPU, so clear it - # before Megatron creates its Gloo DP groups. - default_pg.bound_device_id = None + initialize_process_group(self._should_bind_device_id_for_process_group) diff --git a/src/twinkle/model/transformers/spectral_hybrid_lora.py b/src/twinkle/model/transformers/spectral_hybrid_lora.py index 92faf3043..58781b246 100644 --- a/src/twinkle/model/transformers/spectral_hybrid_lora.py +++ b/src/twinkle/model/transformers/spectral_hybrid_lora.py @@ -117,9 +117,20 @@ def _spectrum_cache_path( return cache_dir / f'spectrum-{digest}.pt' +def resolve_spectral_config_path( + config_value: Optional[str], + output_dir: Path, +) -> Tuple[Path, bool]: + """Return the allocation path and whether an explicitly configured file can be reused.""" + if config_value: + config_path = Path(config_value).expanduser() + return config_path, config_path.is_file() + return Path(output_dir).expanduser() / 'spectral_hybrid_lora_config.json', False + + @torch.no_grad() def compute_spectral_scores( - model: nn.Module, + model: Optional[nn.Module], config: LoraConfig, r: int, *, @@ -127,16 +138,19 @@ def compute_spectral_scores( cache_key: str = '', epsilon: float = 1e-12, log_interval: int = 20, + broadcast: bool = True, ) -> SpectralScores: """Score pretrained modules from singular-value spectra for data-free allocation.""" - targets = select_spectral_targets(model, config) - names = sorted(targets) distributed = dist.is_available() and dist.is_initialized() rank = dist.get_rank() if distributed else 0 scores: Dict[str, float] = {} metrics_by_module: Dict[str, Dict[str, float]] = {} if rank == 0: + if model is None: + raise ValueError('Spectral scoring requires a model on rank 0.') + targets = select_spectral_targets(model, config) + names = sorted(targets) if log_interval > 0: logger.info(f'Spectral Hybrid LoRA scoring: {len(names)} modules, LoRA rank={r}') for index, name in enumerate(names, start=1): @@ -175,7 +189,7 @@ def compute_spectral_scores( f'effective_rank={metrics["effective_rank"]:.1f}, ' f'rank_coverage={metrics["rank_coverage"]:.4f})') - if distributed: + if distributed and broadcast: payload = [(scores, metrics_by_module) if rank == 0 else None] dist.broadcast_object_list(payload, src=0) scores, metrics_by_module = payload[0] diff --git a/tests/transformers/test_spectral_hybrid_lora.py b/tests/transformers/test_spectral_hybrid_lora.py index 27f25cdd2..b58662151 100644 --- a/tests/transformers/test_spectral_hybrid_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -11,6 +11,7 @@ build_spectral_param_groups, compute_spectral_scores, compute_spectral_metrics, + resolve_spectral_config_path, select_spectral_targets, ) @@ -40,6 +41,27 @@ def forward(self, inputs): return inputs +@pytest.mark.parametrize('bind_device,expects_device_id', [(None, True), (lambda _backend: False, False)]) +def test_initialize_process_group_preserves_backend_device_binding(monkeypatch, bind_device, expects_device_id): + import torch.distributed as dist + from twinkle import Platform, torch_util + from twinkle.model.base import initialize_process_group + + calls = [] + monkeypatch.setattr(dist, 'is_initialized', lambda: False) + monkeypatch.setattr(dist, 'init_process_group', lambda **kwargs: calls.append(kwargs)) + monkeypatch.setattr(Platform, 'get_world_size', lambda: 2) + monkeypatch.setattr(Platform, 'get_rank', lambda: 0) + monkeypatch.setattr(Platform, 'get_local_device', lambda: 'cpu') + monkeypatch.setattr(Platform, 'device_backend', lambda: 'nccl') + monkeypatch.setattr(torch_util, 'set_device', lambda: None) + + initialize_process_group(bind_device) + + assert len(calls) == 1 + assert ('device_id' in calls[0]) is expects_device_id + + def test_spectral_metrics_match_weighted_formula(): singular_values = torch.tensor([4.0, 2.0, 1.0, 0.5], dtype=torch.float64) metrics = compute_spectral_metrics(singular_values, r=1) @@ -105,6 +127,51 @@ def record_svdvals(weight): assert all('effective_rank' in scores.metrics[name] for name in scores) +def test_spectral_scores_can_skip_distributed_broadcast(monkeypatch): + from twinkle.model.transformers import spectral_hybrid_lora + + model = TinyDecoder(num_layers=1) + config = LoraConfig(r=2, target_modules=['q_proj']) + monkeypatch.setattr(spectral_hybrid_lora.dist, 'is_available', lambda: True) + monkeypatch.setattr(spectral_hybrid_lora.dist, 'is_initialized', lambda: True) + monkeypatch.setattr(spectral_hybrid_lora.dist, 'get_rank', lambda: 0) + monkeypatch.setattr( + spectral_hybrid_lora.dist, + 'broadcast_object_list', + lambda *_args, **_kwargs: pytest.fail('broadcast should be disabled'), + ) + + scores = compute_spectral_scores(model, config, r=2, log_interval=0, broadcast=False) + + assert set(scores) == {'layers.0.self_attn.q_proj'} + + +def test_config_path_reuses_an_existing_explicit_config(tmp_path): + config_path = tmp_path / 'allocation.json' + config_path.write_text('{}', encoding='utf-8') + + resolved, should_load = resolve_spectral_config_path(str(config_path), tmp_path / 'output') + + assert resolved == config_path + assert should_load is True + + +def test_config_path_computes_when_explicit_config_is_missing(tmp_path): + config_path = tmp_path / 'missing.json' + + resolved, should_load = resolve_spectral_config_path(str(config_path), tmp_path / 'output') + + assert resolved == config_path + assert should_load is False + + +def test_config_path_computes_to_default_when_not_configured(tmp_path): + resolved, should_load = resolve_spectral_config_path(None, tmp_path) + + assert resolved == tmp_path / 'spectral_hybrid_lora_config.json' + assert should_load is False + + def test_allocation_prioritizes_high_scores_within_budget(): scores = {'a': 0.2, 'b': 0.8, 'c': 0.4, 'd': 0.6} counts = {name: 100 for name in scores} From 1ef87769b52555be7d111cb6d969080d67129d1c Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Thu, 13 Aug 2026 17:34:30 +0800 Subject: [PATCH 4/9] feat: support multi-tenant hybrid LoRA training --- src/twinkle/model/__init__.py | 8 +- src/twinkle/model/multi_lora.py | 35 +- .../model/multi_lora_target_parameters.py | 12 +- src/twinkle/model/transformers/__init__.py | 1 + .../model/transformers/hybrid/__init__.py | 4 + .../model/transformers/hybrid/fft_slots.py | 271 ++++++ .../model/transformers/hybrid/model.py | 384 ++++++++ .../spectral_allocation.py} | 25 + .../transformers/multi_lora_transformers.py | 105 ++- .../model/transformers/strategy/accelerate.py | 45 +- .../transformers/strategy/native_fsdp.py | 61 +- src/twinkle/server/config/__init__.py | 4 +- src/twinkle/server/config/application_spec.py | 27 +- src/twinkle/server/model/app.py | 5 +- .../model/backends/transformers_model.py | 20 +- src/twinkle/utils/safetensors.py | 37 +- .../test_multi_lora_target_parameters.py | 38 +- .../transformers/test_spectral_hybrid_lora.py | 833 +++++++++++++++++- 18 files changed, 1814 insertions(+), 101 deletions(-) create mode 100644 src/twinkle/model/transformers/hybrid/__init__.py create mode 100644 src/twinkle/model/transformers/hybrid/fft_slots.py create mode 100644 src/twinkle/model/transformers/hybrid/model.py rename src/twinkle/model/transformers/{spectral_hybrid_lora.py => hybrid/spectral_allocation.py} (91%) diff --git a/src/twinkle/model/__init__.py b/src/twinkle/model/__init__.py index 91401a085..a45c0f8bb 100644 --- a/src/twinkle/model/__init__.py +++ b/src/twinkle/model/__init__.py @@ -6,12 +6,16 @@ if TYPE_CHECKING: from .base import TwinkleModel from .megatron import MegatronModel, MultiLoraMegatronModel - from .transformers import MultiLoraTransformersModel, TransformersModel, TransformersValueModel + from .transformers import (MultiLoraTransformersModel, SpectralHybridTransformersModel, TransformersModel, + TransformersValueModel) else: _import_structure = { 'base': ['TwinkleModel'], - 'transformers': ['TransformersModel', 'MultiLoraTransformersModel', 'TransformersValueModel'], + 'transformers': [ + 'TransformersModel', 'MultiLoraTransformersModel', 'SpectralHybridTransformersModel', + 'TransformersValueModel' + ], 'megatron': ['MegatronModel', 'MultiLoraMegatronModel'], } diff --git a/src/twinkle/model/multi_lora.py b/src/twinkle/model/multi_lora.py index 43cd6108d..f6a80d419 100644 --- a/src/twinkle/model/multi_lora.py +++ b/src/twinkle/model/multi_lora.py @@ -38,6 +38,7 @@ def __init__(self, max_loras=5, max_r=32, max_length: int = 8192): self._active_adapters = [] self.max_length = max_length self.target_parameter_manager = TargetParameterLoraManager(max_loras=max_loras, max_r=max_r) + self.lora_layer_names: List[str] = [] def _get_available_lora(self) -> Optional[LoraTenant]: for _lora in self.loras: @@ -165,8 +166,10 @@ def adapter(self, tenant_adapter_name: str, disable_lora: bool = False): @contextmanager def _disable_lora_context(self, tenant_adapter_name): self.deactivate_adapter() - yield - self.activate_adapter(tenant_adapter_name, call_enable=True) + try: + yield + finally: + self.activate_adapter(tenant_adapter_name, call_enable=True) @contextmanager def save_context(self, tenant_adapter_name: str): @@ -253,6 +256,23 @@ def find_lora(self, adapter_name): else: raise ValueError(f'No lora found for real adapter_name {adapter_name}') + def validate_tenant_target_modules(self, target_modules, target_parameters=None) -> None: + """Ensure every requested target resolves inside the preallocated LoRA layers.""" + if not target_modules and target_parameters: + return + layers = [name for name, layer in self.module.named_modules() if isinstance(layer, LoraLayer)] + if target_modules == 'all-linear' or ( + isinstance(target_modules, (list, set)) and 'all-linear' in target_modules): + return + if isinstance(target_modules, (list, set)): + missing = [target for target in target_modules if not any(name.endswith(target) for name in layers)] + if missing: + raise ValueError(f'LoRA target_modules are outside the preallocated range: {sorted(missing)}') + return + if isinstance(target_modules, str) and any(self.match_target_modules(name, target_modules) for name in layers): + return + raise ValueError(f'LoRA target_modules do not resolve inside the preallocated range: {target_modules!r}') + @staticmethod def match_target_modules( module_name: str, @@ -264,8 +284,7 @@ def match_target_modules( if isinstance(target_modules, (list, set)) and len(target_modules) == 0: return False - if isinstance(target_modules, - (list, set)) and len(target_modules) == 1 and next(iter(target_modules)) == 'all-linear': + if isinstance(target_modules, (list, set)) and 'all-linear' in target_modules: return True if target_modules == 'all-linear': @@ -484,8 +503,9 @@ def patch(self, module_device = next(module.parameters())[1].device low_cpu_mem_usage = module_device.type == 'meta' + base_config = kwargs.get('lora_config', None) for i in range(self.max_loras): - config = kwargs.get('lora_config', None) + config = deepcopy(base_config) if base_config is not None else None if config is None: config = LoraConfig( r=self.max_r, @@ -553,6 +573,11 @@ def _enable_all_lora_grad(_module): _enable_all_lora_grad(module) self.module = module + if not isinstance(module, list): + self.lora_layer_names = [ + name for name, layer in module.named_modules() + if isinstance(layer, Linear) + ] return module def save_initial_weights(self): diff --git a/src/twinkle/model/multi_lora_target_parameters.py b/src/twinkle/model/multi_lora_target_parameters.py index 0eb33eea7..913b4d21e 100644 --- a/src/twinkle/model/multi_lora_target_parameters.py +++ b/src/twinkle/model/multi_lora_target_parameters.py @@ -364,19 +364,25 @@ def parameters_for_tenant(self, tenant_adapter_name: str) -> list[nn.Parameter]: return parameters def named_slot_parameters(self, tenant_adapter_name: str) -> Iterator[tuple[str, nn.Parameter]]: - slot_name = self.tenant_to_slot[tenant_adapter_name] + slot_name = self.tenant_to_slot.get(tenant_adapter_name) + if slot_name is None: + return for wrapper in self.wrappers: yield from wrapper.named_slot_parameters(slot_name) def get_state_dict(self, tenant_adapter_name: str) -> dict[str, torch.Tensor]: - slot_name = self.tenant_to_slot[tenant_adapter_name] + slot_name = self.tenant_to_slot.get(tenant_adapter_name) + if slot_name is None: + return {} state_dict = {} for wrapper in self.wrappers: state_dict.update(wrapper.get_state_dict(slot_name)) return state_dict def set_state_dict(self, tenant_adapter_name: str, state_dict: dict[str, torch.Tensor]) -> set[str]: - slot_name = self.tenant_to_slot[tenant_adapter_name] + slot_name = self.tenant_to_slot.get(tenant_adapter_name) + if slot_name is None: + return set() consumed_keys = set() for wrapper in self.wrappers: consumed_keys.update(wrapper.set_state_dict(slot_name, state_dict)) diff --git a/src/twinkle/model/transformers/__init__.py b/src/twinkle/model/transformers/__init__.py index cd3775ace..e72c003df 100644 --- a/src/twinkle/model/transformers/__init__.py +++ b/src/twinkle/model/transformers/__init__.py @@ -1,4 +1,5 @@ # Copyright (c) ModelScope Contributors. All rights reserved. from .multi_lora_transformers import MultiLoraTransformersModel +from .hybrid import SpectralHybridTransformersModel from .transformers import TransformersModel from .value_model import TransformersValueModel diff --git a/src/twinkle/model/transformers/hybrid/__init__.py b/src/twinkle/model/transformers/hybrid/__init__.py new file mode 100644 index 000000000..d084feda5 --- /dev/null +++ b/src/twinkle/model/transformers/hybrid/__init__.py @@ -0,0 +1,4 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from .model import SpectralHybridTransformersModel + +__all__ = ['SpectralHybridTransformersModel'] diff --git a/src/twinkle/model/transformers/hybrid/fft_slots.py b/src/twinkle/model/transformers/hybrid/fft_slots.py new file mode 100644 index 000000000..cb77fd7ad --- /dev/null +++ b/src/twinkle/model/transformers/hybrid/fft_slots.py @@ -0,0 +1,271 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +from typing import Dict, List + +import torch +from peft.tuners.lora import LoraLayer +from peft.utils import ModulesToSaveWrapper +from torch import nn + +from twinkle.model.multi_lora import MultiLora + + +class HybridFftSlots: + """Manage the full-module FFT slots used by Hybrid adapters. + + ``MultiLora`` remains responsible for LoRA slots. This class owns only + the server-selected FFT modules and derives the matching ``fft_N`` slot + from a tenant's existing LoRA slot index. + """ + + def __init__(self, multi_lora: MultiLora, s_fft: List[str]) -> None: + if not s_fft: + raise ValueError('Hybrid requires at least one S_FFT module.') + if len(set(s_fft)) != len(s_fft): + raise ValueError('Hybrid S_FFT contains the same module more than once.') + normalized_s_fft = [self._canonical_module_name(name) for name in s_fft] + if len(set(normalized_s_fft)) != len(normalized_s_fft): + raise ValueError('Hybrid S_FFT aliases resolve to the same layer.') + self.multi_lora = multi_lora + self.s_fft = normalized_s_fft + self.hybrid_adapters: set[str] = set() + self.allocated_to_layer_name: Dict[str, str] = {} + self.allocated_to_wrapper_name: Dict[str, str] = {} + + @property + def module(self): + return self.multi_lora.module + + @staticmethod + def _canonical_module_name(name: str) -> str: + prefix = 'base_model.model.' + return name[len(prefix):] if name.startswith(prefix) else name + + def install_fft_slots(self) -> None: + """Install PEFT ``ModulesToSaveWrapper`` slots before DDP/FSDP wrapping.""" + if isinstance(self.module, list): + raise NotImplementedError('Hybrid FFT slots currently require the Transformers backend.') + named_modules = dict(self.module.named_modules()) + resolved_layer_names = set() + for allocated_name in self.s_fft: + matches = [ + (name, layer) for name, layer in named_modules.items() + if self._canonical_module_name(name) == allocated_name + and isinstance(layer, (LoraLayer, nn.Linear)) + ] + if len(matches) != 1: + raise ValueError( + f'Hybrid S_FFT module {allocated_name!r} resolved to {len(matches)} layers.') + layer_name, layer = matches[0] + if layer_name in resolved_layer_names: + raise ValueError( + f'Hybrid S_FFT aliases resolve to the same layer {layer_name!r}.') + resolved_layer_names.add(layer_name) + if isinstance(layer, LoraLayer): + original_module = layer.base_layer + wrapper_name = f'{layer_name}.base_layer' + else: + original_module = layer + wrapper_name = layer_name + if any(parameter.is_meta for parameter in original_module.parameters()): + raise ValueError('Hybrid FFT slots require materialized base weights.') + wrapper = ModulesToSaveWrapper(original_module, 'fft_0') + for slot in range(1, self.multi_lora.max_loras): + wrapper.update(f'fft_{slot}') + wrapper.set_adapter([]) + for parameter in wrapper.modules_to_save.parameters(): + parameter.requires_grad_(True) + if isinstance(layer, LoraLayer): + layer.base_layer = wrapper + else: + parent_name, _, child_name = layer_name.rpartition('.') + parent = self.module.get_submodule(parent_name) if parent_name else self.module + setattr(parent, child_name, wrapper) + self.allocated_to_layer_name[allocated_name] = layer_name + self.allocated_to_wrapper_name[allocated_name] = wrapper_name + + def is_hybrid(self, adapter_name: str) -> bool: + return adapter_name in self.hybrid_adapters + + def register_adapter(self, adapter_name: str) -> None: + self.hybrid_adapters.add(adapter_name) + + def unregister_adapter(self, adapter_name: str) -> None: + self.hybrid_adapters.discard(adapter_name) + + def _tenant(self, adapter_name: str): + return self.multi_lora.find_lora_by_tenant(adapter_name) + + def _fft_adapter_name(self, adapter_name: str) -> str: + return f'fft_{self._tenant(adapter_name).index}' + + def _get_fft_wrapper(self, allocated_name: str) -> ModulesToSaveWrapper: + return self.module.get_submodule(self.allocated_to_wrapper_name[allocated_name]) + + def _iter_fft_wrappers(self): + return [self._get_fft_wrapper(name) for name in self.s_fft] + + def activate_fft_slot(self, adapter_name: str) -> None: + fft_adapter_name = self._fft_adapter_name(adapter_name) if self.is_hybrid(adapter_name) else None + for wrapper in self._iter_fft_wrappers(): + wrapper.set_adapter(fft_adapter_name if fft_adapter_name is not None else []) + + def deactivate_fft_slots(self) -> None: + for wrapper in self._iter_fft_wrappers(): + wrapper.set_adapter([]) + + def resolve_lora_targets(self, target_modules) -> List[str]: + """Resolve tenant LoRA targets while reserving S_FFT for full tuning.""" + lora_layers = self.multi_lora.lora_layer_names + fft_layers = set(self.allocated_to_layer_name.values()) + if target_modules is None: + selected = lora_layers + else: + selected = [ + name for name in lora_layers + if self.multi_lora.match_target_modules(name, target_modules) + ] + return sorted(name for name in selected if name not in fft_layers) + + def _tenant_lora_layer_names(self, adapter_name: str) -> List[str]: + tenant = self._tenant(adapter_name) + fft_layers = set(self.allocated_to_layer_name.values()) + return [ + name for name in self.multi_lora.lora_layer_names + if name not in fft_layers + and self.multi_lora.match_target_modules(name, tenant.tenant_config.target_modules) + ] + + @staticmethod + def _iter_module_tensors(module: nn.Module): + for name, parameter in module.named_parameters(): + yield name, parameter, True + for name, buffer in module.named_buffers(): + yield name, buffer, False + + def _iter_fft_slot_tensors(self, adapter_name: str): + fft_adapter_name = self._fft_adapter_name(adapter_name) + for allocated_name in self.s_fft: + wrapper_name = self.allocated_to_wrapper_name[allocated_name] + slot_module = self._get_fft_wrapper(allocated_name).modules_to_save[fft_adapter_name] + for tensor_name, tensor, is_parameter in self._iter_module_tensors(slot_module): + yield allocated_name, wrapper_name, fft_adapter_name, tensor_name, tensor, is_parameter + + def reset_adapter_slot(self, adapter_name: str) -> None: + for allocated_name, _, _, tensor_name, target, _ in self._iter_fft_slot_tensors(adapter_name): + wrapper = self._get_fft_wrapper(allocated_name) + original_tensors = { + name: tensor + for name, tensor, _ in self._iter_module_tensors(wrapper.original_module) + } + self.multi_lora._write_param_tensor( + target, self.multi_lora._read_param_tensor(original_tensors[tensor_name])) + + @staticmethod + def _checkpoint_key(allocated_name: str, parameter_name: str) -> str: + return f'base_model.model.{allocated_name}.{parameter_name}' + + def get_fft_state_dict(self, adapter_name: str) -> Dict[str, torch.Tensor]: + if not self.is_hybrid(adapter_name): + return {} + state = {} + for allocated_name, _, _, tensor_name, tensor, _ in self._iter_fft_slot_tensors(adapter_name): + state[self._checkpoint_key(allocated_name, tensor_name)] = ( + self.multi_lora._read_param_tensor(tensor).detach().clone()) + return state + + def set_fft_state_dict(self, adapter_name: str, state_dict: Dict[str, torch.Tensor]) -> None: + if not self.is_hybrid(adapter_name): + return + for allocated_name, _, _, tensor_name, tensor, _ in self._iter_fft_slot_tensors(adapter_name): + key = self._checkpoint_key(allocated_name, tensor_name) + if key not in state_dict: + raise ValueError(f'Hybrid training state is missing {key!r}.') + self.multi_lora._write_param_tensor(tensor, state_dict[key]) + + def named_fft_parameters(self, adapter_name: str): + if not self.is_hybrid(adapter_name): + return [] + result = [] + for _, wrapper_name, fft_adapter_name, tensor_name, tensor, is_parameter in self._iter_fft_slot_tensors( + adapter_name): + if is_parameter: + result.append((f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}', tensor)) + return result + + @staticmethod + def _normalize_base_state_key(name: str) -> str: + prefix = 'base_model.model.' + if name.startswith(prefix): + name = name[len(prefix):] + name = name.replace('.base_layer.original_module.', '.') + name = name.replace('.original_module.', '.') + return name.replace('.base_layer.', '.') + + def iter_merged_state_dict(self, adapter_name: str, full_state_dict: Dict[str, torch.Tensor]): + """Yield a non-destructively merged plain Transformers state dict.""" + if not self.is_hybrid(adapter_name): + raise ValueError(f'Adapter {adapter_name!r} is not a Hybrid adapter.') + tenant = self._tenant(adapter_name) + slot = tenant.adapter_name + replacements: Dict[str, torch.Tensor] = {} + for layer_name in self._tenant_lora_layer_names(adapter_name): + base_key = f'{layer_name}.base_layer.weight' + a_key = f'{layer_name}.lora_A.{slot}.weight' + b_key = f'{layer_name}.lora_B.{slot}.weight' + missing = [key for key in (base_key, a_key, b_key) if key not in full_state_dict] + if missing: + raise ValueError(f'Cannot export Hybrid module {layer_name!r}; missing {missing}.') + base = full_state_dict[base_key] + rank = tenant.tenant_config.r + a = full_state_dict[a_key][:rank, :] + b = full_state_dict[b_key][:, :rank] + scaling = tenant.tenant_config.lora_alpha / ( + rank**0.5 if tenant.tenant_config.use_rslora else rank) + delta = b.to(torch.float32) @ a.to(torch.float32) + if getattr(tenant.tenant_config, 'fan_in_fan_out', False): + delta = delta.transpose(0, 1) + replacements[base_key] = base + delta.to(dtype=base.dtype) * scaling + + for allocated_name, wrapper_name, fft_adapter_name, tensor_name, _, _ in self._iter_fft_slot_tensors( + adapter_name): + base_key = f'{wrapper_name}.original_module.{tensor_name}' + fft_key = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}' + if base_key not in full_state_dict or fft_key not in full_state_dict: + raise ValueError(f'Cannot export Hybrid FFT layer {allocated_name!r}.') + replacements[base_key] = full_state_dict[fft_key] + + emitted = set() + for name, value in full_state_dict.items(): + if ('.lora_' in name or '.modules_to_save.' in name + or self.multi_lora._is_target_parameter_lora_name(name)): + continue + output_name = self._normalize_base_state_key(name) + if output_name in emitted: + continue + emitted.add(output_name) + yield output_name, replacements.get(name, value).detach().cpu() + + def build_merged_state_dict(self, adapter_name: str, full_state_dict: Dict[str, torch.Tensor]): + return dict(self.iter_merged_state_dict(adapter_name, full_state_dict)) + + def build_training_state_dict(self, adapter_name: str, full_state_dict: Dict[str, torch.Tensor]): + """Extract lossless LoRA and FFT state for a Hybrid tenant.""" + if not self.is_hybrid(adapter_name): + raise ValueError(f'Adapter {adapter_name!r} is not a Hybrid adapter.') + tenant = self._tenant(adapter_name) + state = {} + for layer_name in self._tenant_lora_layer_names(adapter_name): + for kind in ('A', 'B'): + source = f'{layer_name}.lora_{kind}.{tenant.adapter_name}.weight' + if source not in full_state_dict: + raise ValueError(f'Hybrid training state is missing {source!r}.') + value = self.multi_lora._slice_rank_tensor( + source, full_state_dict[source], tenant.tenant_config.r) + state[source.replace(f'.{tenant.adapter_name}.', '.')] = value.detach().cpu() + for allocated_name, wrapper_name, fft_adapter_name, tensor_name, _, _ in self._iter_fft_slot_tensors( + adapter_name): + source = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}' + if source not in full_state_dict: + raise ValueError(f'Hybrid training state is missing {source!r}.') + state[self._checkpoint_key(allocated_name, tensor_name)] = full_state_dict[source].detach().cpu() + return state diff --git a/src/twinkle/model/transformers/hybrid/model.py b/src/twinkle/model/transformers/hybrid/model.py new file mode 100644 index 000000000..65dd3aaa5 --- /dev/null +++ b/src/twinkle/model/transformers/hybrid/model.py @@ -0,0 +1,384 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import json +import os +import shutil +from contextlib import contextmanager +from typing import Any, Dict, List, Optional, Type, Union + +import torch +import torch.distributed as dist +from peft import PeftConfig +from safetensors.torch import load_file, save_file +from torch.optim import Optimizer + +from twinkle import Platform, remote_class, remote_function +from twinkle.utils.safetensors import StreamingSafetensorSaver + +from ..multi_lora_transformers import MultiLoraTransformersModel +from .fft_slots import HybridFftSlots +from .spectral_allocation import load_spectral_allocation + + +_HYBRID_CONFIG_FIELDS = ( + 'r', + 'lora_alpha', + 'lora_dropout', + 'use_rslora', + 'fan_in_fan_out', + 'use_dora', + 'bias', + 'rank_pattern', + 'alpha_pattern', + 'target_modules', + 'modules_to_save', +) +HYBRID_ADAPTER_MODE = 'hybrid' + + +@remote_class() +class SpectralHybridTransformersModel(MultiLoraTransformersModel): + """Transformers MultiLoRA service extended with server-owned FFT slots.""" + + def __init__(self, hybrid: Dict[str, Any], memory_efficient_init: bool = False, **kwargs): + if memory_efficient_init: + raise ValueError( + 'Spectral Hybrid does not support memory_efficient_init because FFT slots require materialized ' + 'base weights before FSDP wrapping.') + config = dict(hybrid or {}) + allocation_path = config.get('allocation_path') + if not allocation_path: + raise ValueError('Spectral Hybrid requires allocation_path.') + s_fft = load_spectral_allocation(allocation_path) + self.default_lr_lora = float(config.get('default_lr_lora', 2.5e-5)) + self.default_lr_fft = float(config.get('default_lr_fft', 1e-6)) + super().__init__(memory_efficient_init=False, **kwargs) + self.fft_slots = HybridFftSlots(self.multi_adapter, s_fft) + self.fft_slots.install_fft_slots() + + @contextmanager + def _adapter_context(self, adapter_name: str, disable_lora: bool = False): + with super()._adapter_context(adapter_name, disable_lora=disable_lora) as slot_name: + if disable_lora: + self.fft_slots.deactivate_fft_slots() + else: + self.fft_slots.activate_fft_slot(adapter_name) + try: + yield slot_name + finally: + self.fft_slots.deactivate_fft_slots() + + @remote_function() + def add_adapter_to_model(self, adapter_name: str, config_or_dir: Union[PeftConfig, str], **kwargs): + adapter_mode = kwargs.pop('adapter_mode', 'lora') + if adapter_mode == 'lora': + return super().add_adapter_to_model(adapter_name, config_or_dir, **kwargs) + if adapter_mode != HYBRID_ADAPTER_MODE: + raise ValueError( + f'Unsupported adapter_mode {adapter_mode!r}; expected "lora" or {HYBRID_ADAPTER_MODE!r}.') + config = self._copy_lora_config(config_or_dir) + if config.modules_to_save: + raise ValueError('Hybrid modules_to_save is controlled by the server allocation.') + if getattr(config, 'target_parameters', None): + raise ValueError('Hybrid target_parameters is not supported.') + config.target_modules = set(self.fft_slots.resolve_lora_targets(config.target_modules)) + config.modules_to_save = list(self.fft_slots.s_fft) + self._register_adapter(adapter_name, config, **kwargs) + self.fft_slots.register_adapter(adapter_name) + + def _create_param_group(self, adapter_name: str, lr: float = 1e-5, weight_decay: float = 0.01, **kwargs): + if not self.fft_slots.is_hybrid(adapter_name): + return super()._create_param_group( + adapter_name=adapter_name, lr=lr, weight_decay=weight_decay, **kwargs) + params = self._get_trainable_parameters(adapter_name) + fft_token = '.modules_to_save.fft_' + lora_names = [name for name in params if fft_token not in name] + fft_names = [name for name in params if fft_token in name] + groups = [] + if lora_names: + groups.append({ + 'params': [params[name] for name in lora_names], + 'param_names': lora_names, + 'lr': kwargs.get('lr_lora', self.default_lr_lora), + 'weight_decay': weight_decay, + }) + if fft_names: + groups.append({ + 'params': [params[name] for name in fft_names], + 'param_names': fft_names, + 'lr': kwargs.get('lr_fft', self.default_lr_fft), + 'weight_decay': weight_decay, + }) + if not groups: + raise ValueError(f'Spectral Hybrid adapter {adapter_name!r} has no trainable parameters.') + return groups + + @remote_function() + def set_optimizer(self, optimizer_cls: Union[Type[Optimizer], str], **kwargs): + adapter_name = kwargs.get('adapter_name') + if not self.fft_slots.is_hybrid(adapter_name): + return super().set_optimizer(optimizer_cls, **kwargs) + if 'params' not in kwargs: + lr_lora = kwargs.pop('lr_lora', kwargs.get('lr', self.default_lr_lora)) + lr_fft = kwargs.pop('lr_fft', self.default_lr_fft) + kwargs['params'] = self._create_param_group( + adapter_name, + weight_decay=kwargs.get('weight_decay', 0.01), + lr_lora=lr_lora, + lr_fft=lr_fft, + ) + return super().set_optimizer(optimizer_cls, **kwargs) + + def _get_trainable_parameters(self, adapter_name): + params = super()._get_trainable_parameters(adapter_name) + if not self.fft_slots.is_hybrid(adapter_name): + return params + known_parameter_ids = {id(parameter) for parameter in params.values()} + for name, parameter in self.fft_slots.named_fft_parameters(adapter_name): + if id(parameter) not in known_parameter_ids: + params[name] = parameter + known_parameter_ids.add(id(parameter)) + return params + + @remote_function(collect='first') + def get_state_dict(self, **kwargs): + adapter_name = kwargs.get('adapter_name') + self._check_adapter_valid(adapter_name) + state = self.multi_adapter.get_state_dict(adapter_name) + if self.fft_slots.is_hybrid(adapter_name): + state.update(self.fft_slots.get_fft_state_dict(adapter_name)) + return state + + def _validate_hybrid_training_checkpoint_boundary(self, adapter_name: str) -> None: + """Require a quiescent optimizer-step boundary for a lossless checkpoint.""" + optimizer_group = self.optimizer_group[adapter_name] + if optimizer_group.optimizer is None: + raise ValueError('Spectral Hybrid optimizer must be configured before save_optimizer=True.') + for group in optimizer_group.optimizer.param_groups: + if len(group.get('param_names', [])) != len(group['params']): + raise ValueError( + 'Spectral Hybrid lossless checkpoints require optimizer param_names for every parameter.') + train_status = optimizer_group.train_status + if train_status.loss_value is not None or train_status.num_tokens != 0: + raise ValueError( + 'Spectral Hybrid training state can only be saved after the optimizer step and zero_grad.') + if any(parameter.grad is not None for parameter in self._get_trainable_parameters(adapter_name).values()): + raise ValueError( + 'Spectral Hybrid training state can only be saved after zero_grad cleared accumulated gradients.') + if (optimizer_group.cur_step > 0 and optimizer_group.gradient_accumulation_steps > 1 + and not optimizer_group.do_grad_sync()): + raise ValueError( + 'Spectral Hybrid training state cannot be saved in the middle of gradient accumulation.') + + @staticmethod + def _class_identity(instance) -> Optional[str]: + if instance is None: + return None + cls = instance.__class__ + return f'{cls.__module__}.{cls.__qualname__}' + + @staticmethod + def _normalize_hybrid_config(config) -> dict: + raw = config.to_dict() if hasattr(config, 'to_dict') else dict(config) + normalized = {} + for field in _HYBRID_CONFIG_FIELDS: + value = raw.get(field) + if field in ('target_modules', 'modules_to_save'): + value = sorted(value or []) + elif field in ('rank_pattern', 'alpha_pattern'): + value = dict(sorted((value or {}).items())) + normalized[field] = value + return normalized + + @staticmethod + def _optimizer_param_names(optimizer) -> List[List[str]]: + return [list(group.get('param_names', [])) for group in optimizer.param_groups] + + def _write_hybrid_training_manifest(self, training_dir: str, adapter_name: str) -> None: + if not Platform.is_master(): + return + optimizer_group = self.optimizer_group[adapter_name] + trainer_state_path = os.path.join(training_dir, 'trainer_state.json') + with open(trainer_state_path, encoding='utf-8') as handle: + trainer_state = json.load(handle) + trainer_state.update({ + 'checkpoint_boundary': 'optimizer_step', + 'optimizer_class': self._class_identity(optimizer_group.optimizer), + 'optimizer_param_names': self._optimizer_param_names(optimizer_group.optimizer), + 'scheduler_class': self._class_identity(optimizer_group.lr_scheduler), + 'has_scaler': optimizer_group.scaler is not None, + }) + with open(trainer_state_path, 'w', encoding='utf-8') as handle: + json.dump(trainer_state, handle, ensure_ascii=False, indent=2, sort_keys=True) + handle.write('\n') + + def _validate_hybrid_resume_state(self, training_dir: str, adapter_name: str, saved_config: dict, + trainer_state: dict) -> None: + optimizer_group = self.optimizer_group[adapter_name] + current_config = self._normalize_hybrid_config(optimizer_group.adapter_config) + checkpoint_config = self._normalize_hybrid_config(saved_config) + differences = { + field: (checkpoint_config[field], current_config[field]) + for field in _HYBRID_CONFIG_FIELDS if checkpoint_config[field] != current_config[field] + } + if differences: + raise ValueError(f'Spectral Hybrid adapter config does not match checkpoint: {differences}') + if trainer_state.get('checkpoint_boundary') != 'optimizer_step': + raise ValueError('Spectral Hybrid checkpoint was not saved at a supported optimizer-step boundary.') + if optimizer_group.optimizer is None: + raise ValueError('Spectral Hybrid optimizer must be configured before resuming training.') + optimizer_path = os.path.join(training_dir, 'optimizer.pt') + if not os.path.isfile(optimizer_path): + raise ValueError('Spectral Hybrid training state is missing optimizer.pt.') + if trainer_state.get('optimizer_class') != self._class_identity(optimizer_group.optimizer): + raise ValueError('Spectral Hybrid optimizer class does not match the checkpoint.') + if trainer_state.get('optimizer_param_names') != self._optimizer_param_names(optimizer_group.optimizer): + raise ValueError('Spectral Hybrid optimizer parameter groups do not match the checkpoint.') + saved_scheduler = trainer_state.get('scheduler_class') + if saved_scheduler != self._class_identity(optimizer_group.lr_scheduler): + raise ValueError('Spectral Hybrid scheduler configuration does not match the checkpoint.') + if saved_scheduler and not os.path.isfile(os.path.join(training_dir, 'scheduler.pt')): + raise ValueError('Spectral Hybrid training state is missing scheduler.pt.') + saved_scaler = bool(trainer_state.get('has_scaler')) + if saved_scaler != (optimizer_group.scaler is not None): + raise ValueError('Spectral Hybrid grad scaler configuration does not match the checkpoint.') + if saved_scaler and not os.path.isfile(os.path.join(training_dir, 'scaler.pt')): + raise ValueError('Spectral Hybrid training state is missing scaler.pt.') + rank = dist.get_rank() if dist.is_initialized() else 0 + rank_rng_path = os.path.join(training_dir, f'rng_state_rank_{rank}.pt') + if not os.path.isfile(rank_rng_path): + raise ValueError(f'Spectral Hybrid training state is missing rank RNG state: {rank_rng_path}') + + def _save_spectral_hybrid(self, name, output_dir: Optional[str], interval: int, adapter_name: str, **kwargs): + optimizer_group = self.optimizer_group[adapter_name] + if name is None: + name = f'checkpoint-step-{optimizer_group.cur_step}' + output_dir = output_dir or 'output' + checkpoint_dir = os.path.join(output_dir, name) + if optimizer_group.cur_step % interval != 0: + return None + if kwargs.get('save_optimizer', False): + self._validate_hybrid_training_checkpoint_boundary(adapter_name) + training_dir = os.path.join(checkpoint_dir, 'twinkle_training_state') + if Platform.is_master() and os.path.isdir(training_dir): + shutil.rmtree(training_dir) + + full_state = self.strategy.get_full_state_dict(self.model) + saver = StreamingSafetensorSaver( + checkpoint_dir, + max_shard_size=kwargs.get('max_shard_size', '5GB'), + save_rank='master', + ) + if Platform.is_master(): + for key, value in self.fft_slots.iter_merged_state_dict(adapter_name, full_state): + saver.add_tensor(key, value) + saver.finalize() + + model = self.strategy.unwrap_model(self.model) + if Platform.is_master(): + self.hf_config.save_pretrained(checkpoint_dir) + generation_config = getattr(model, 'generation_config', None) + if generation_config is not None: + generation_config.save_pretrained(checkpoint_dir) + else: + generation_config_path = os.path.join(checkpoint_dir, 'generation_config.json') + if os.path.exists(generation_config_path): + os.unlink(generation_config_path) + self._save_tokenizer(checkpoint_dir, adapter_name=adapter_name) + + if kwargs.get('save_optimizer', False): + if Platform.is_master(): + adapter_state = self.fft_slots.build_training_state_dict(adapter_name, full_state) + os.makedirs(training_dir, exist_ok=True) + optimizer_group.adapter_config.save_pretrained(training_dir) + config_path = os.path.join(training_dir, 'adapter_config.json') + with open(config_path, encoding='utf-8') as handle: + adapter_config = json.load(handle) + adapter_config['twinkle_adapter_mode'] = HYBRID_ADAPTER_MODE + with open(config_path, 'w', encoding='utf-8') as handle: + json.dump(adapter_config, handle, ensure_ascii=False, indent=2, sort_keys=True) + handle.write('\n') + save_file( + {key: value.contiguous() for key, value in adapter_state.items()}, + os.path.join(training_dir, 'adapter_model.safetensors'), + ) + if dist.is_initialized(): + dist.barrier() + self._save_training_state( + training_dir, + adapter_name=adapter_name, + consumed_train_samples=kwargs.get('consumed_train_samples', 0), + ) + self._write_hybrid_training_manifest(training_dir, adapter_name) + rank = dist.get_rank() if dist.is_initialized() else 0 + torch.save(self._get_training_rng_state(), os.path.join(training_dir, f'rng_state_rank_{rank}.pt')) + if dist.is_initialized(): + dist.barrier() + return checkpoint_dir + + @remote_function(collect='first') + def save(self, name, output_dir: Optional[str] = None, interval=1, **kwargs): + adapter_name = kwargs.get('adapter_name') + self._check_adapter_valid(adapter_name) + if not self.fft_slots.is_hybrid(adapter_name): + return super().save(name, output_dir, interval, **kwargs) + checkpoint_dir = self._save_spectral_hybrid(name, output_dir, interval, adapter_name, **kwargs) + if dist.is_initialized(): + dist.barrier() + return checkpoint_dir + + def _resume_spectral_hybrid(self, checkpoint_dir: str, adapter_name: str, resume_only_model: bool): + training_dir = os.path.join(checkpoint_dir, 'twinkle_training_state') + adapter_path = os.path.join(training_dir, 'adapter_model.safetensors') + trainer_state_path = os.path.join(training_dir, 'trainer_state.json') + if not os.path.isfile(adapter_path) or not os.path.isfile(trainer_state_path): + raise ValueError( + 'Cannot resume Spectral Hybrid training from a merged-only checkpoint. ' + 'Save the checkpoint with save_optimizer=True to create twinkle_training_state.') + adapter_config_path = os.path.join(training_dir, 'adapter_config.json') + if not os.path.isfile(adapter_config_path): + raise ValueError('Spectral Hybrid training state is missing adapter_config.json.') + with open(adapter_config_path, encoding='utf-8') as handle: + saved_config = json.load(handle) + if saved_config.get('twinkle_adapter_mode') != HYBRID_ADAPTER_MODE: + raise ValueError('Checkpoint training state is not a Spectral Hybrid adapter.') + with open(trainer_state_path, encoding='utf-8') as handle: + trainer_state = json.load(handle) + if not resume_only_model: + self._validate_hybrid_resume_state(training_dir, adapter_name, saved_config, trainer_state) + else: + current_config = self._normalize_hybrid_config(self.optimizer_group[adapter_name].adapter_config) + checkpoint_config = self._normalize_hybrid_config(saved_config) + if current_config != checkpoint_config: + raise ValueError('Spectral Hybrid adapter config does not match checkpoint.') + + adapter_state = load_file(adapter_path, device='cpu') + self.multi_adapter.set_state_dict(adapter_name, adapter_state) + self.fft_slots.set_fft_state_dict(adapter_name, adapter_state) + if not resume_only_model: + trainer_state = self._restore_training_state(training_dir, adapter_name=adapter_name) + rank = dist.get_rank() if dist.is_initialized() else 0 + self._load_rng_state(os.path.join(training_dir, f'rng_state_rank_{rank}.pt')) + return { + 'cur_step': trainer_state['cur_step'], + 'consumed_train_samples': trainer_state['consumed_train_samples'], + 'gradient_accumulation_steps': trainer_state['gradient_accumulation_steps'], + } + + @remote_function(dispatch='all', collect='first', sync=True) + def resume_from_checkpoint(self, checkpoint_dir, *, resume_only_model=False, **kwargs): + adapter_name = kwargs.get('adapter_name', '') + self._check_adapter_valid(adapter_name) + if not self.fft_slots.is_hybrid(adapter_name): + return super().resume_from_checkpoint( + checkpoint_dir, resume_only_model=resume_only_model, **kwargs) + result = self._resume_spectral_hybrid(checkpoint_dir, adapter_name, resume_only_model) + if dist.is_initialized(): + dist.barrier() + return result + + @remote_function() + def remove_adapter(self, adapter_name: str): + if self.fft_slots.is_hybrid(adapter_name): + self.fft_slots.reset_adapter_slot(adapter_name) + self.fft_slots.unregister_adapter(adapter_name) + return super().remove_adapter(adapter_name) diff --git a/src/twinkle/model/transformers/spectral_hybrid_lora.py b/src/twinkle/model/transformers/hybrid/spectral_allocation.py similarity index 91% rename from src/twinkle/model/transformers/spectral_hybrid_lora.py rename to src/twinkle/model/transformers/hybrid/spectral_allocation.py index 58781b246..e54dbb32d 100644 --- a/src/twinkle/model/transformers/spectral_hybrid_lora.py +++ b/src/twinkle/model/transformers/hybrid/spectral_allocation.py @@ -1,6 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import copy import hashlib +import json import os import re import torch @@ -28,6 +29,30 @@ _LAYER_RE = re.compile(r'\blayers\.(\d+)\.') +def load_spectral_allocation(config_path: str | Path) -> List[str]: + """Load the server-owned FFT module allocation.""" + path = Path(config_path).expanduser() + if not path.is_file(): + raise ValueError(f'Spectral Hybrid allocation file does not exist: {path}') + with path.open(encoding='utf-8') as handle: + raw = json.load(handle) + if not isinstance(raw, dict): + raise ValueError('Spectral Hybrid allocation JSON must contain an object.') + + def _modules(keys: Tuple[str, ...]) -> List[str]: + value = next((raw[key] for key in keys if key in raw), []) + if isinstance(value, str): + value = [value] + if not isinstance(value, (list, tuple, set)) or not all(isinstance(item, str) for item in value): + raise ValueError('Spectral Hybrid s_fft must be a list of module names.') + return sorted(value) + + s_fft = _modules(('s_fft', 'S_FFT', 'modules_to_save')) + if not s_fft: + raise ValueError('Spectral Hybrid allocation requires at least one S_FFT module.') + return s_fft + + class SpectralScores(dict[str, float]): """Spectral scores with per-module metric details.""" diff --git a/src/twinkle/model/transformers/multi_lora_transformers.py b/src/twinkle/model/transformers/multi_lora_transformers.py index ea53930de..00e1b409b 100644 --- a/src/twinkle/model/transformers/multi_lora_transformers.py +++ b/src/twinkle/model/transformers/multi_lora_transformers.py @@ -1,5 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import os +from contextlib import contextmanager +from copy import deepcopy import torch.distributed as dist import transformers from peft import LoraConfig, PeftConfig, PeftModel, load_peft_weights @@ -39,6 +41,7 @@ def __init__( max_r: int = 32, max_length: int = 8192, target_modules: Union[List[str], str] = 'all-linear', + preallocated_lora_modules: Optional[Union[List[str], str]] = None, **kwargs): os.environ['TOKENIZERS_PARALLELISM'] = 'true' self._try_init_process_group() @@ -81,6 +84,9 @@ def __init__( self.sp_strategy = None # Initialize expert parallel attributes (required by set_optimizer in TransformersModel) self.optimizer_group: Dict[str, OptimizerGroup] = {} + if preallocated_lora_modules is not None: + target_modules = preallocated_lora_modules + self.preallocated_lora_modules = target_modules self.multi_adapter = MultiLora(max_loras=max_loras, max_r=max_r, max_length=max_length) self.model.gradient_checkpointing_enable() self.model = self.multi_adapter.patch(self.model, target_modules=target_modules, lora_config=self.lora_config) @@ -132,6 +138,12 @@ def _ensure_target_parameter_lora_installed(self, config: LoraConfig) -> None: # self._maybe_apply_expert_parallel() # 各rank广播之前不能对moe层进行分片, 没有实际权重时不能分片 self.multi_adapter.patch_target_parameters(self.model, target_parameters) + @contextmanager + def _adapter_context(self, adapter_name: str, disable_lora: bool = False): + """Activate one LoRA tenant for a model operation.""" + with self.multi_adapter.adapter(adapter_name, disable_lora=disable_lora) as slot_name: + yield slot_name + @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) def forward(self, *, inputs: Union[InputFeature, List[InputFeature], Trajectory, List[Trajectory]], **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) @@ -145,7 +157,7 @@ def forward(self, *, inputs: Union[InputFeature, List[InputFeature], Trajectory, inputs = [inputs] inputs = optimizer_config.template.batch_encode(inputs) # noqa self.multi_adapter.check_length(inputs) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): return super().forward(inputs=inputs, **kwargs) @remote_function(dispatch='slice_dp', collect=collect_tensor_dict) @@ -163,25 +175,25 @@ def forward_only(self, *, inputs: Union[InputFeature, List[InputFeature], List[T inputs = [inputs] inputs = optimizer_config.template.batch_encode(inputs) # noqa self.multi_adapter.check_length(inputs) - with self.multi_adapter.adapter(adapter_name, disable_lora=disable_lora): + with self._adapter_context(adapter_name, disable_lora=disable_lora): return super().forward_only(inputs=inputs, **kwargs) @remote_function(collect='mean') def calculate_loss(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): return super().calculate_loss(**kwargs) @remote_function() def backward(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): super().backward(**kwargs) @remote_function() def clip_grad_norm(self, max_grad_norm: float = 1.0, norm_type=2, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): return super().clip_grad_norm(max_grad_norm, norm_type=norm_type, **kwargs) def _create_param_group(self, adapter_name: str, lr: float = 1e-5, weight_decay: float = 0.01, **kwargs): @@ -190,19 +202,19 @@ def _create_param_group(self, adapter_name: str, lr: float = 1e-5, weight_decay: @remote_function() def step(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): super().step(**kwargs) @remote_function() def zero_grad(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): super().zero_grad(**kwargs) @remote_function() def lr_step(self, **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): super().lr_step(**kwargs) @remote_function() @@ -213,35 +225,52 @@ def set_loss(self, loss_cls: Union[Type[Loss], str], **kwargs): @remote_function() def set_optimizer(self, optimizer_cls: Union[Type[Optimizer], str], **kwargs): self._check_adapter_valid(kwargs.get('adapter_name')) - with self.multi_adapter.adapter(kwargs.get('adapter_name')): + with self._adapter_context(kwargs.get('adapter_name')): super().set_optimizer(optimizer_cls, **kwargs) @remote_function() def add_adapter_to_model(self, adapter_name: str, config_or_dir: Union[PeftConfig, str], **kwargs): - # prevent opening requires_grad of the base model - # prevent loading malicious code - assert not isinstance( - config_or_dir, str - ), 'config_or_dir does not support str, because loading config from modelhub may causing unexpected behavior' - assert isinstance(config_or_dir, LoraConfig), 'config_or_dir must be a LoraConfig instance' - config_or_dir = self.strategy.prepare_adapter_config( - config_or_dir, + adapter_mode = kwargs.pop('adapter_mode', 'lora') + if adapter_mode != 'lora': + raise ValueError(f'MultiLoraTransformersModel only supports LoRA adapters, got {adapter_mode!r}.') + config = self._copy_lora_config(config_or_dir) + if config.modules_to_save: + raise ValueError('modules_to_save is not supported for a multi-tenant LoRA adapter.') + self.multi_adapter.validate_tenant_target_modules( + config.target_modules, getattr(config, 'target_parameters', None)) + self._register_adapter(adapter_name, config, **kwargs) + + @staticmethod + def _copy_lora_config(config_or_dir: Union[PeftConfig, str]) -> LoraConfig: + """Validate and copy a client-provided LoRA config before normalizing it.""" + if isinstance(config_or_dir, str): + raise ValueError('Loading an adapter config from a model path or hub is not supported.') + if not isinstance(config_or_dir, LoraConfig): + raise TypeError('config_or_dir must be a LoraConfig instance.') + return deepcopy(config_or_dir) + + def _register_adapter(self, adapter_name: str, config: LoraConfig, **kwargs) -> None: + """Register an already validated LoRA configuration in a free slot.""" + config = self.strategy.prepare_adapter_config( + config, enable_ep=getattr(self, '_enable_expert_parallel', False), ) # Limit the max peft version in pyproject.toml, in case any newer version opens some untested module grad. - config_or_dir.modules_to_save = None - config_or_dir.bias = 'none' - config_or_dir.init_lora_weights = False - config_or_dir.modules_to_save = None - config_or_dir.trainable_token_indices = None + config.bias = 'none' + config.init_lora_weights = False + config.trainable_token_indices = None self.optimizer_group[adapter_name] = self._construct_default_optimizer_group() self.optimizer_group[adapter_name].adapter_name = adapter_name - self.optimizer_group[adapter_name].adapter_config = config_or_dir + self.optimizer_group[adapter_name].adapter_config = config _gas_default = kwargs.get('gradient_accumulation_steps', 1) self.optimizer_group[adapter_name].gradient_accumulation_steps = _gas_default self._default_tokenizer = self.optimizer_group[adapter_name].template.processor - self._ensure_target_parameter_lora_installed(config_or_dir) - self.multi_adapter.acquire_lora(tenant_adapter_name=adapter_name, config=config_or_dir) + self._ensure_target_parameter_lora_installed(config) + try: + self.multi_adapter.acquire_lora(tenant_adapter_name=adapter_name, config=config) + except Exception: + self.optimizer_group.pop(adapter_name, None) + raise @remote_function() def set_lr_scheduler(self, scheduler_cls: Union[Type[LRScheduler], str], **kwargs): @@ -323,18 +352,28 @@ def remove_adapter(self, adapter_name: str): self.multi_adapter.release_lora(adapter_name) def _get_nb_trainable_parameters(self, adapter_name, model): - with self.multi_adapter.adapter(adapter_name): + with self._adapter_context(adapter_name): return self.multi_adapter.get_nb_trainable_parameters(adapter_name) def _get_trainable_parameters_example(self, adapter_name, model): - with self.multi_adapter.adapter(adapter_name): + with self._adapter_context(adapter_name): return self.multi_adapter.get_trainable_parameters_example(adapter_name) def _get_trainable_parameters(self, adapter_name): - with self.multi_adapter.adapter(adapter_name) as real_adapter_name: - params = super()._get_trainable_parameters(real_adapter_name) - # Note: experts have registered LoraWrapper as a submodule, so its internal LoRA parameters - # are already captured automatically. Duplicating parameter capture here will cause - # optimizer errors due to duplicate keys. - # params.update(self.multi_adapter.get_target_parameter_trainable_parameters(adapter_name)) + with self._adapter_context(adapter_name) as real_adapter_name: + tenant = self.multi_adapter.find_lora_by_tenant(adapter_name) + pattern = f'.{real_adapter_name}.' + params = {} + model = self.strategy.unwrap_model(self.model) + for name, parameter in model.named_parameters(): + if not parameter.requires_grad: + continue + if pattern in name and '.lora_' in name: + if self.multi_adapter.match_target_modules(name, tenant.tenant_config.target_modules): + params[name] = parameter + known_parameter_ids = {id(parameter) for parameter in params.values()} + for name, parameter in self.multi_adapter.target_parameter_manager.named_slot_parameters(adapter_name): + if id(parameter) not in known_parameter_ids: + params[name] = parameter + known_parameter_ids.add(id(parameter)) return params diff --git a/src/twinkle/model/transformers/strategy/accelerate.py b/src/twinkle/model/transformers/strategy/accelerate.py index fdd808355..a8e77669b 100644 --- a/src/twinkle/model/transformers/strategy/accelerate.py +++ b/src/twinkle/model/transformers/strategy/accelerate.py @@ -161,6 +161,16 @@ def _prepare_fsdp2_sd_options(self): broadcast_from_rank0=getattr(fsdp_plugin.state_dict_config, 'rank0_only', False), ) + @staticmethod + def _prepare_full_optimizer_state_dict_options(*, for_load: bool): + from torch.distributed.checkpoint.state_dict import StateDictOptions + + return StateDictOptions( + full_state_dict=True, + cpu_offload=not for_load, + broadcast_from_rank0=for_load, + ) + def needs_wrapped_optimizer_state(self) -> bool: fsdp_plugin = self._get_fsdp_plugin() return fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2 @@ -171,7 +181,11 @@ def save_optimizer_checkpoint(self, model, optimizer, output_path: str): if fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2: from torch.distributed.checkpoint.state_dict import get_optimizer_state_dict - optim_state = get_optimizer_state_dict(model, optimizer, options=self._prepare_fsdp2_sd_options()) + optim_state = get_optimizer_state_dict( + model, + optimizer, + options=self._prepare_full_optimizer_state_dict_options(for_load=False), + ) if self.accelerator.process_index == 0: torch.save(optim_state, output_path) return @@ -185,11 +199,15 @@ def load_optimizer_checkpoint(self, model, optimizer, input_path: str): if fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2: from torch.distributed.checkpoint.state_dict import set_optimizer_state_dict - optim_state = None - rank0_only = getattr(fsdp_plugin.optim_state_dict_config, 'rank0_only', False) - if self.accelerator.process_index == 0 or not rank0_only: - optim_state = torch.load(input_path, weights_only=True) - set_optimizer_state_dict(model, optimizer, optim_state, options=self._prepare_fsdp2_sd_options()) + optim_state = {} + if self.accelerator.process_index == 0: + optim_state = torch.load(input_path, map_location='cpu', weights_only=True) + set_optimizer_state_dict( + model, + optimizer, + optim_state, + options=self._prepare_full_optimizer_state_dict_options(for_load=True), + ) return optimizer.load_state_dict(torch.load(input_path, map_location='cpu', weights_only=False)) @@ -197,10 +215,21 @@ def load_optimizer_checkpoint(self, model, optimizer, input_path: str): def get_full_state_dict(self, model) -> dict: """Collect full state dict.""" from twinkle.utils import torch_util + fsdp_plugin = self._get_fsdp_plugin() + if fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2: + from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict + + # A merged checkpoint is a deployment artifact, so it must always + # materialize global, CPU-offloaded parameters on rank 0 regardless + # of the training plugin's normal sharded checkpoint preference. + return get_model_state_dict( + model, + options=StateDictOptions(full_state_dict=True, cpu_offload=True), + ) unwrapped = self.unwrap_model(model) state_dict = {} - for name, param in unwrapped.named_parameters(): - local = torch_util.to_local_tensor(param) + for name, value in unwrapped.state_dict().items(): + local = torch_util.to_local_tensor(value) state_dict[name] = local.cpu() del local return state_dict diff --git a/src/twinkle/model/transformers/strategy/native_fsdp.py b/src/twinkle/model/transformers/strategy/native_fsdp.py index 92dbb8d62..cce92bd03 100644 --- a/src/twinkle/model/transformers/strategy/native_fsdp.py +++ b/src/twinkle/model/transformers/strategy/native_fsdp.py @@ -318,32 +318,45 @@ def get_full_state_dict(self, model) -> dict: the local expert shards across the EP group to reconstruct the full expert tensor (all num_experts on dim-0). """ + if self.device_mesh is not None: + from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict + + ep_mesh = self.ep_fsdp_device_mesh + ep_world_size = ep_mesh['ep'].size() if ep_mesh is not None else 1 + if ep_world_size <= 1: + # FSDP2 parameters must be gathered through the state-dict API; + # CPU offload makes the result rank0-only. + return get_model_state_dict( + model, + options=StateDictOptions(full_state_dict=True, cpu_offload=True), + ) + + # EP experts are independently sharded across the EP dimension, so + # every EP rank must retain its FSDP-gathered expert block until the + # second all-gather below has reconstructed the original tensor. + state_dict = get_model_state_dict( + model, + options=StateDictOptions(full_state_dict=True, cpu_offload=False), + ) + unwrapped = self.unwrap_model(model) + ep_expert_names = _detect_ep_expert_names(unwrapped) + ep_group = ep_mesh['ep'].get_group() + result = {} + for name, value in state_dict.items(): + if name in ep_expert_names: + local_full = value.contiguous().to(Platform.get_local_device()) + gathered = [torch.empty_like(local_full) for _ in range(ep_world_size)] + dist.all_gather(gathered, local_full, group=ep_group) + value = torch.cat(gathered, dim=_ep_expert_state_dict_gather_dim(name)) + if Platform.is_master(): + result[name] = value.cpu() + return result unwrapped = self.unwrap_model(model) state_dict = {} - - ep_fsdp_mesh = self.ep_fsdp_device_mesh - ep_group = None - ep_world_size = 1 - if ep_fsdp_mesh is not None: - ep_group = ep_fsdp_mesh['ep'].get_group() - ep_world_size = ep_fsdp_mesh['ep'].size() - - ep_expert_names = _detect_ep_expert_names(unwrapped) if ep_world_size > 1 else set() - - for name, param in unwrapped.named_parameters(): - local_full = torch_util.to_local_tensor(param) - - if name in ep_expert_names and ep_world_size > 1 and ep_group is not None: - local_full = local_full.contiguous().to(Platform.get_local_device()) - gathered = [torch.empty_like(local_full) for _ in range(ep_world_size)] - dist.all_gather(gathered, local_full, group=ep_group) - local_full = torch.cat(gathered, dim=_ep_expert_state_dict_gather_dim(name)) - state_dict[name] = local_full.cpu() - del gathered, local_full - else: - state_dict[name] = local_full.cpu() - del local_full - + for name, value in unwrapped.state_dict().items(): + local_full = torch_util.to_local_tensor(value) + state_dict[name] = local_full.cpu() + del local_full return state_dict def get_adapter_state_dict(self, model, adapter_name: str) -> dict: diff --git a/src/twinkle/server/config/__init__.py b/src/twinkle/server/config/__init__.py index dfdd5176d..463099ece 100644 --- a/src/twinkle/server/config/__init__.py +++ b/src/twinkle/server/config/__init__.py @@ -1,7 +1,8 @@ # Copyright (c) ModelScope Contributors. All rights reserved. """Server configuration package — aggregate root and per-deployment specs.""" -from .application_spec import ApplicationSpec, HttpOptions, ModelArgs, ProcessorArgs, SamplerArgs, ServerArgs +from .application_spec import (ApplicationSpec, HttpOptions, ModelArgs, ProcessorArgs, SamplerArgs, ServerArgs, + SpectralHybridArgs) from .persistence import PersistenceConfig from .server_config import ServerConfig from .telemetry import TelemetryConfig @@ -14,6 +15,7 @@ 'ProcessorArgs', 'SamplerArgs', 'ServerArgs', + 'SpectralHybridArgs', 'ServerConfig', 'TelemetryConfig', ] diff --git a/src/twinkle/server/config/application_spec.py b/src/twinkle/server/config/application_spec.py index 51015245c..20db9dd84 100644 --- a/src/twinkle/server/config/application_spec.py +++ b/src/twinkle/server/config/application_spec.py @@ -11,7 +11,8 @@ """ from __future__ import annotations -from pydantic import BaseModel, ConfigDict, Field, model_validator +import math +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator from typing import Any, Literal from twinkle.server.utils.task_queue.config import TaskQueueConfig @@ -41,6 +42,21 @@ class HttpOptions(BaseModel): # ---------- per-deployment args schemas ------------------------------------ # +class SpectralHybridArgs(_ArgsBase): + """Strict server-owned Spectral Hybrid allocation and optimizer defaults.""" + + allocation_path: str = Field(min_length=1) + default_lr_lora: float = Field(default=2.5e-5, gt=0) + default_lr_fft: float = Field(default=1.0e-6, gt=0) + + @field_validator('default_lr_lora', 'default_lr_fft') + @classmethod + def _finite_learning_rate(cls, value: float) -> float: + if not math.isfinite(value): + raise ValueError('learning rate must be finite') + return value + + class ModelArgs(_ArgsBase): """Args for the ``model`` deployment. @@ -57,7 +73,16 @@ class ModelArgs(_ArgsBase): adapter_config: dict[str, Any] | None = None queue_config: TaskQueueConfig = Field(default_factory=TaskQueueConfig) max_loras: int = 5 + max_r: int = Field(default=32, gt=0) max_length: int | None = None + preallocated_lora_modules: str | list[str] = 'all-linear' + hybrid: SpectralHybridArgs | None = None + + @model_validator(mode='after') + def _validate_spectral_backend(self): + if self.hybrid is not None and self.backend != 'transformers': + raise ValueError('hybrid is only supported by the transformers backend') + return self class SamplerArgs(_ArgsBase): diff --git a/src/twinkle/server/model/app.py b/src/twinkle/server/model/app.py index bae7b2a93..03b096adc 100644 --- a/src/twinkle/server/model/app.py +++ b/src/twinkle/server/model/app.py @@ -35,8 +35,11 @@ def _make_mock_model(kw: dict[str, Any]) -> Any: def _make_transformers_model(kw: dict[str, Any]) -> Any: - from .backends.transformers_model import TwinkleCompatTransformersModel + from .backends.transformers_model import (TwinkleCompatSpectralHybridTransformersModel, + TwinkleCompatTransformersModel) + if kw.get('hybrid'): + return TwinkleCompatSpectralHybridTransformersModel(**kw) return TwinkleCompatTransformersModel(**kw) diff --git a/src/twinkle/server/model/backends/transformers_model.py b/src/twinkle/server/model/backends/transformers_model.py index e7677b619..b488fae49 100644 --- a/src/twinkle/server/model/backends/transformers_model.py +++ b/src/twinkle/server/model/backends/transformers_model.py @@ -12,16 +12,15 @@ from twinkle import remote_class, remote_function from twinkle.data_format import InputFeature, Trajectory from twinkle.infra import collect_tensor_dict -from twinkle.model import MultiLoraTransformersModel +from twinkle.model import MultiLoraTransformersModel, SpectralHybridTransformersModel from twinkle.server.common.datum import datum_to_input_feature, extract_rl_features_for_loss from twinkle.server.model.backends.common import (TwinkleCompatModelBase, clean_metrics, collect_forward_backward_results, to_cpu_safe_output) from twinkle.utils.nccl_safe import nccl_safe -@remote_class() -class TwinkleCompatTransformersModel(MultiLoraTransformersModel, TwinkleCompatModelBase): - """Unified wrapper around MultiLoraTransformersModel. +class _TwinkleCompatTransformersMixin: + """Tinker and Twinkle API compatibility shared by Transformers backends. Handles both: - Tinker-compat I/O (Datum / TensorData) via /tinker/* endpoints. @@ -110,3 +109,16 @@ def forward_backward(self, *, inputs: InputFeature | list[InputFeature] | Trajec def ping(self) -> bool: """Lightweight liveness probe for watchdog health checks.""" return True + + +@remote_class() +class TwinkleCompatTransformersModel(_TwinkleCompatTransformersMixin, MultiLoraTransformersModel, + TwinkleCompatModelBase): + """Unified API wrapper for the pure MultiLoRA Transformers model.""" + + +@remote_class() +class TwinkleCompatSpectralHybridTransformersModel(_TwinkleCompatTransformersMixin, + SpectralHybridTransformersModel, + TwinkleCompatModelBase): + """Unified API wrapper for the Spectral Hybrid Transformers model.""" diff --git a/src/twinkle/utils/safetensors.py b/src/twinkle/utils/safetensors.py index 42619c54d..9d28647b9 100644 --- a/src/twinkle/utils/safetensors.py +++ b/src/twinkle/utils/safetensors.py @@ -1,5 +1,6 @@ import json import os +from pathlib import Path from functools import partial from typing import Literal @@ -89,11 +90,18 @@ def __init__( ) -> None: self.save_dir = save_dir if isinstance(max_shard_size, str): - if max_shard_size.endswith('GB'): - max_shard_size = int(max_shard_size[:-2]) - else: + import re + match = re.fullmatch(r'\s*(\d+(?:\.\d+)?)\s*(B|KB|MB|GB)\s*', max_shard_size.upper()) + if match is None: raise ValueError(f'Invalid max_shard_size: {max_shard_size}') - self.max_shard_size = max_shard_size * 1000**3 + multipliers = {'B': 1, 'KB': 1000, 'MB': 1000**2, 'GB': 1000**3} + max_shard_size = int(float(match.group(1)) * multipliers[match.group(2)]) + elif isinstance(max_shard_size, int): + # Preserve the historical API: integer values are expressed in GB. + max_shard_size *= 1000**3 + else: + raise TypeError('max_shard_size must be a byte count or a size string such as "5GB".') + self.max_shard_size = max_shard_size self.current_shard = {} self.current_shard_size = 0 self.total_size = 0 @@ -103,6 +111,12 @@ def __init__( self.is_peft_format = is_peft_format if self.is_save_rank: os.makedirs(save_dir, exist_ok=True) + self._previous_model_files = set() + if self.is_save_rank and not self.is_peft_format: + self._previous_model_files = set(Path(save_dir).glob('model*.safetensors')) + index_path = Path(save_dir) / 'model.safetensors.index.json' + if index_path.exists(): + self._previous_model_files.add(index_path) def add_tensor(self, name, tensor): if not self.is_save_rank: @@ -112,7 +126,9 @@ def add_tensor(self, name, tensor): and not self.is_peft_format): self._save_current_shard() - self.current_shard[name] = tensor.cpu().contiguous() + # Break shared storage (for example tied input/output embeddings) so a + # standard Transformers state dict remains valid for safetensors. + self.current_shard[name] = tensor.detach().cpu().contiguous().clone() self.current_shard_size += tensor_size def _save_current_shard(self, shard_filename: str = None): @@ -161,6 +177,17 @@ def finalize(self): self._save_index(updated_weight_map) + if total_shards == 1: + current_files = {Path(self.save_dir) / 'model.safetensors'} + else: + current_files = { + Path(self.save_dir) / f'model-{index:05d}-of-{total_shards:05d}.safetensors' + for index in range(1, total_shards + 1) + } + current_files.add(Path(self.save_dir) / 'model.safetensors.index.json') + for stale_path in self._previous_model_files - current_files: + stale_path.unlink(missing_ok=True) + def _save_index(self, weight_map): index = {'metadata': {'total_size': self.total_size}, 'weight_map': weight_map} diff --git a/tests/model/test_multi_lora_target_parameters.py b/tests/model/test_multi_lora_target_parameters.py index b28ef6b36..e6a68e79a 100644 --- a/tests/model/test_multi_lora_target_parameters.py +++ b/tests/model/test_multi_lora_target_parameters.py @@ -9,6 +9,7 @@ print(f"sys.path: {sys.path}") + class FakePackedExperts(nn.Module): def __init__(self, num_experts=2, hidden=4, intermediate=6, *, is_transposed=False): @@ -133,7 +134,7 @@ def test_multilora_releases_target_parameter_slot_to_initial_weights(): else: assert torch.count_nonzero(param.detach()) == 0 -# Note: PEFT (Parameter-Efficient Fine-Tuning) does not natively support +# Note: PEFT (Parameter-Efficient Fine-Tuning) does not natively support # installing multiple LoRA slots on target parameters. # def test_target_parameter_state_dict_loads_with_peft(): # from twinkle.model.multi_lora_target_parameters import TargetParameterLoraManager @@ -260,10 +261,31 @@ def test_multilora_transformers_installs_target_parameters_once(): else: raise AssertionError("different target_parameters should be rejected") -# Run in the local environment. -if __name__ == "__main__": - assert test_peft_target_parameter_key_shapes_for_3d_experts() == True - assert test_target_parameter_multi_lora_updates_only_active_adapter() == True - assert test_multilora_releases_target_parameter_slot_to_initial_weights() == True - assert test_multilora_state_dict_round_trips_target_parameters() == True - assert test_multilora_transformers_installs_target_parameters_once() == True \ No newline at end of file + +def test_multilora_transformers_optimizer_includes_target_parameter_slots(): + from twinkle.model.transformers.multi_lora_transformers import MultiLoraTransformersModel + + model = FakeModel() + manager = _make_multilora_for_target_parameters(model) + manager.acquire_lora("adapter_a", _make_target_cfg(r=2)) + + instance = object.__new__(MultiLoraTransformersModel) + instance.multi_adapter = manager + instance.strategy = type("Strategy", (), {"unwrap_model": lambda _self, inner: inner})() + instance.__dict__["model"] = model + + selected = instance._get_trainable_parameters("adapter_a") + expected = dict(manager.target_parameter_manager.named_slot_parameters("adapter_a")) + + assert expected + assert set(selected) == set(expected) + assert {id(value) for value in selected.values()} == {id(value) for value in expected.values()} + + +def test_target_parameter_only_adapter_does_not_require_linear_targets(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + manager.module = FakeModel() + + manager.validate_tenant_target_modules([], target_parameters=_make_target_cfg().target_parameters) diff --git a/tests/transformers/test_spectral_hybrid_lora.py b/tests/transformers/test_spectral_hybrid_lora.py index b58662151..06edab19b 100644 --- a/tests/transformers/test_spectral_hybrid_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -1,16 +1,20 @@ +from contextlib import contextmanager + import pytest import torch from peft import LoraConfig, PeftModel, get_peft_model from peft.utils import get_peft_model_state_dict from torch import nn -from twinkle.model.transformers.spectral_hybrid_lora import ( +from twinkle.model.transformers.hybrid.fft_slots import HybridFftSlots +from twinkle.model.transformers.hybrid.spectral_allocation import ( CANDIDATE_TYPES, allocate_spectral_modules, build_spectral_lora_config, build_spectral_param_groups, compute_spectral_scores, compute_spectral_metrics, + load_spectral_allocation, resolve_spectral_config_path, select_spectral_targets, ) @@ -41,6 +45,38 @@ def forward(self, inputs): return inputs +def _fft_slot_module(hybrid, allocated_name, slot=0): + wrapper = hybrid._get_fft_wrapper(allocated_name) + return wrapper.modules_to_save[f'fft_{slot}'] + + +@contextmanager +def _adapter(manager, hybrid, adapter_name, disable_lora=False): + with manager.adapter(adapter_name, disable_lora=disable_lora): + if disable_lora: + hybrid.deactivate_fft_slots() + else: + hybrid.activate_fft_slot(adapter_name) + try: + yield + finally: + hybrid.deactivate_fft_slots() + + +def _install_hybrid(manager, model, s_fft): + hybrid = HybridFftSlots(manager, s_fft) + manager.module = model + hybrid.install_fft_slots() + return hybrid + + +def _register_hybrid(manager, hybrid, adapter_name, config): + config.target_modules = set(hybrid.resolve_lora_targets(config.target_modules)) + config.modules_to_save = list(hybrid.s_fft) + manager.acquire_lora(adapter_name, config) + hybrid.register_adapter(adapter_name) + + @pytest.mark.parametrize('bind_device,expects_device_id', [(None, True), (lambda _backend: False, False)]) def test_initialize_process_group_preserves_backend_device_binding(monkeypatch, bind_device, expects_device_id): import torch.distributed as dist @@ -128,15 +164,15 @@ def record_svdvals(weight): def test_spectral_scores_can_skip_distributed_broadcast(monkeypatch): - from twinkle.model.transformers import spectral_hybrid_lora + from twinkle.model.transformers.hybrid import spectral_allocation model = TinyDecoder(num_layers=1) config = LoraConfig(r=2, target_modules=['q_proj']) - monkeypatch.setattr(spectral_hybrid_lora.dist, 'is_available', lambda: True) - monkeypatch.setattr(spectral_hybrid_lora.dist, 'is_initialized', lambda: True) - monkeypatch.setattr(spectral_hybrid_lora.dist, 'get_rank', lambda: 0) + monkeypatch.setattr(spectral_allocation.dist, 'is_available', lambda: True) + monkeypatch.setattr(spectral_allocation.dist, 'is_initialized', lambda: True) + monkeypatch.setattr(spectral_allocation.dist, 'get_rank', lambda: 0) monkeypatch.setattr( - spectral_hybrid_lora.dist, + spectral_allocation.dist, 'broadcast_object_list', lambda *_args, **_kwargs: pytest.fail('broadcast should be disabled'), ) @@ -247,6 +283,99 @@ def test_strategy_adapter_state_includes_full_modules(strategy_cls): assert any('.modules_to_save.default.' in name for name in state) +def test_accelerate_fsdp_optimizer_uses_full_rank0_state_options(): + from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy + + save_options = AccelerateStrategy._prepare_full_optimizer_state_dict_options(for_load=False) + load_options = AccelerateStrategy._prepare_full_optimizer_state_dict_options(for_load=True) + + assert save_options.full_state_dict is True + assert save_options.cpu_offload is True + assert save_options.broadcast_from_rank0 is False + assert load_options.full_state_dict is True + assert load_options.cpu_offload is False + assert load_options.broadcast_from_rank0 is True + + +def test_accelerate_fsdp_optimizer_save_and_load_ignore_sharded_plugin_options(tmp_path, monkeypatch): + from types import SimpleNamespace + from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy + import torch.distributed.checkpoint.state_dict as state_dict_api + + strategy = object.__new__(AccelerateStrategy) + strategy.accelerator = SimpleNamespace(process_index=0) + strategy._get_fsdp_plugin = lambda: SimpleNamespace(fsdp_version=2) + calls = {} + + def fake_get(_model, _optimizer, *, options): + calls['save'] = options + return {'state': {}, 'param_groups': []} + + def fake_set(_model, _optimizer, state, *, options): + calls['load'] = (state, options) + + monkeypatch.setattr(state_dict_api, 'get_optimizer_state_dict', fake_get) + monkeypatch.setattr(state_dict_api, 'set_optimizer_state_dict', fake_set) + checkpoint = tmp_path / 'optimizer.pt' + strategy.save_optimizer_checkpoint(object(), object(), str(checkpoint)) + strategy.load_optimizer_checkpoint(object(), object(), str(checkpoint)) + + assert calls['save'].full_state_dict is True + assert calls['save'].cpu_offload is True + loaded_state, load_options = calls['load'] + assert loaded_state == {'state': {}, 'param_groups': []} + assert load_options.full_state_dict is True + assert load_options.broadcast_from_rank0 is True + + +def test_native_fsdp_full_state_reconstructs_ep_experts(monkeypatch): + import torch.distributed.checkpoint.state_dict as state_dict_api + from twinkle import Platform + from twinkle.model.transformers.strategy import native_fsdp + from twinkle.model.transformers.strategy.native_fsdp import NativeFSDPStrategy + + class EpDimension: + + @staticmethod + def size(): + return 2 + + @staticmethod + def get_group(): + return 'ep-group' + + strategy = object.__new__(NativeFSDPStrategy) + strategy.device_mesh = object() + strategy.ep_fsdp_device_mesh = {'ep': EpDimension()} + strategy.unwrap_model = lambda model: model + captured_options = [] + + def get_model_state_dict(_model, *, options): + captured_options.append(options) + return { + 'experts.weight': torch.tensor([[1.0]]), + 'dense.weight': torch.tensor([[3.0]]), + } + + def all_gather(output, value, *, group): + assert group == 'ep-group' + output[0].copy_(value) + output[1].copy_(value + 1) + + monkeypatch.setattr(state_dict_api, 'get_model_state_dict', get_model_state_dict) + monkeypatch.setattr(native_fsdp, '_detect_ep_expert_names', lambda _model: {'experts.weight'}) + monkeypatch.setattr(native_fsdp.dist, 'all_gather', all_gather) + monkeypatch.setattr(Platform, 'get_local_device', lambda: 'cpu') + monkeypatch.setattr(Platform, 'is_master', lambda: True) + + state = strategy.get_full_state_dict(object()) + + assert captured_options[0].full_state_dict is True + assert captured_options[0].cpu_offload is False + assert torch.equal(state['experts.weight'], torch.tensor([[1.0], [2.0]])) + assert torch.equal(state['dense.weight'], torch.tensor([[3.0]])) + + def test_twinkle_checkpoint_normalization_round_trips_full_modules(tmp_path): from safetensors.torch import save_file from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy @@ -309,3 +438,695 @@ def test_trainable_parameter_filter_includes_full_modules(adapter_name): assert set(selected) == expected assert any('.modules_to_save.' in name for name in selected) + + +def test_load_fixed_server_allocation_only_requires_fft(tmp_path): + allocation = tmp_path / 'allocation.json' + allocation.write_text( + '{"method":"spectral_hybrid",' + '"s_fft":["layers.0.self_attn.q_proj"]}', + encoding='utf-8', + ) + + s_fft = load_spectral_allocation(allocation) + + assert s_fft == ['layers.0.self_attn.q_proj'] + + allocation.write_text('{"s_lora":["ignored"]}', encoding='utf-8') + with pytest.raises(ValueError, match='at least one S_FFT'): + load_spectral_allocation(allocation) + + +def test_server_hybrid_config_is_strict(): + from pydantic import ValidationError + from twinkle.server.config.application_spec import ModelArgs + + base = { + 'model_id': 'tiny', + 'device_group': {}, + 'device_mesh': {}, + 'backend': 'transformers', + } + config = ModelArgs(**base, hybrid={'allocation_path': '/shared/allocation.json'}) + assert config.hybrid.default_lr_fft == pytest.approx(1e-6) + + with pytest.raises(ValidationError, match='extra_forbidden'): + ModelArgs(**base, spectral_hybrid={'allocation_path': '/shared/allocation.json'}) + with pytest.raises(ValidationError, match='extra_forbidden'): + ModelArgs(**base, hybrid={ + 'allocation_path': '/shared/allocation.json', + 'default_lr_fFt': 1e-5, + }) + with pytest.raises(ValidationError, match='greater_than'): + ModelArgs(**base, hybrid={ + 'allocation_path': '/shared/allocation.json', + 'default_lr_fft': -1.0, + }) + with pytest.raises(ValidationError, match='only supported by the transformers backend'): + ModelArgs(**{**base, 'backend': 'megatron'}, hybrid={ + 'allocation_path': '/shared/allocation.json', + }) + + +def test_lora_config_copy_is_explicit_and_non_mutating(): + from twinkle.model.transformers.multi_lora_transformers import MultiLoraTransformersModel + + original = LoraConfig(r=2, lora_alpha=4, target_modules=['q_proj']) + copied = MultiLoraTransformersModel._copy_lora_config(original) + + assert copied is not original + assert copied.to_dict() == original.to_dict() + with pytest.raises(ValueError, match='model path or hub'): + MultiLoraTransformersModel._copy_lora_config('/adapter/path') + with pytest.raises(TypeError, match='LoraConfig'): + MultiLoraTransformersModel._copy_lora_config(object()) + + +def test_hybrid_rejects_client_owned_fft_and_target_parameter_config(): + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + model = object.__new__(SpectralHybridTransformersModel) + modules_to_save = LoraConfig( + r=2, + target_modules=['q_proj'], + modules_to_save=['k_proj'], + ) + with pytest.raises(ValueError, match='modules_to_save is controlled by the server'): + model.add_adapter_to_model('tenant', modules_to_save, adapter_mode='hybrid') + + target_parameters = LoraConfig(r=2, target_modules=['q_proj']) + target_parameters.target_parameters = ['weight'] + with pytest.raises(ValueError, match='target_parameters is not supported'): + model.add_adapter_to_model('tenant', target_parameters, adapter_mode='hybrid') + + +def test_transformers_server_selects_hybrid_model_only_when_configured(monkeypatch): + from twinkle.server.model import app + from twinkle.server.model.backends import transformers_model + + class Regular: + + def __init__(self, **kwargs): + self.kwargs = kwargs + + class Hybrid(Regular): + pass + + monkeypatch.setattr(transformers_model, 'TwinkleCompatTransformersModel', Regular) + monkeypatch.setattr(transformers_model, 'TwinkleCompatSpectralHybridTransformersModel', Hybrid) + + assert type(app._make_transformers_model({'model_id': 'tiny'})) is Regular + assert type(app._make_transformers_model({ + 'model_id': 'tiny', + 'hybrid': {'allocation_path': '/shared/allocation.json'}, + })) is Hybrid + + +def test_allocation_rejects_duplicate_modules(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + with pytest.raises(ValueError, match='same module more than once'): + HybridFftSlots( + manager, + ['layers.0.self_attn.q_proj', 'layers.0.self_attn.q_proj'], + ) + + manager = MultiLora(max_loras=1, max_r=4) + manager.patch(TinyDecoder(num_layers=1), target_modules='all-linear') + hybrid = HybridFftSlots(manager, ['q_proj']) + with pytest.raises(ValueError, match='resolved to 0 layers'): + hybrid.install_fft_slots() + + +def test_allocation_rejects_aliases_for_the_same_module(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + manager.patch(TinyDecoder(num_layers=1), target_modules='all-linear') + with pytest.raises(ValueError, match='aliases resolve to the same layer'): + HybridFftSlots( + manager, + ['layers.0.self_attn.q_proj', 'base_model.model.layers.0.self_attn.q_proj'], + ) + + +def test_fft_modules_do_not_need_lora_preallocation(): + from peft.utils import ModulesToSaveWrapper + from twinkle.model.multi_lora import MultiLora + + base = TinyDecoder(num_layers=1) + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(base, target_modules=['down_proj']) + hybrid = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + config = LoraConfig(r=2, lora_alpha=4, target_modules=['layers.0.mlp.down_proj']) + _register_hybrid(manager, hybrid, 'hybrid', config) + + fft_wrapper = hybrid._get_fft_wrapper('layers.0.self_attn.q_proj') + fft_layer = _fft_slot_module(hybrid, 'layers.0.self_attn.q_proj') + assert isinstance(fft_wrapper, ModulesToSaveWrapper) + assert isinstance(fft_layer, nn.Linear) + assert not any('_twinkle_fft' in name for name, _ in model.named_parameters()) + with torch.no_grad(): + fft_layer.weight.add_(0.1) + inputs = torch.randn(2, 8) + with _adapter(manager, hybrid, 'hybrid'): + expected = model(inputs).detach() + full_state = {name: value.detach().clone() for name, value in model.state_dict().items()} + merged_state = hybrid.build_merged_state_dict('hybrid', full_state) + deployed = TinyDecoder(num_layers=1) + deployed.load_state_dict(merged_state) + assert torch.allclose(deployed(inputs), expected, atol=1e-5) + + +def test_hybrid_lora_targets_follow_tenant_config_but_exclude_fft(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(TinyDecoder(num_layers=1), target_modules=['down_proj']) + hybrid = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + + targets = hybrid.resolve_lora_targets('all-linear') + + assert len(targets) == 1 + assert any(name.endswith('layers.0.mlp.down_proj') for name in targets) + assert not any(name.endswith('layers.0.self_attn.q_proj') for name in targets) + + only_fft = hybrid.resolve_lora_targets(['q_proj']) + assert only_fft == [] + + +def test_fft_layer_never_stacks_a_lora_delta(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(TinyDecoder(num_layers=1), target_modules='all-linear') + hybrid = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + config = LoraConfig(r=2, lora_alpha=4, target_modules='all-linear', init_lora_weights=False) + _register_hybrid(manager, hybrid, 'hybrid', config) + tenant = manager.find_lora_by_tenant('hybrid') + q_proj = model.get_submodule(hybrid.allocated_to_layer_name['layers.0.self_attn.q_proj']) + inputs = torch.randn(2, 8) + + with _adapter(manager, hybrid, 'hybrid'): + before = model(inputs).detach() + with torch.no_grad(): + q_proj.lora_A[tenant.adapter_name].weight[:2].fill_(10.0) + q_proj.lora_B[tenant.adapter_name].weight[:, :2].fill_(10.0) + with _adapter(manager, hybrid, 'hybrid'): + after = model(inputs).detach() + + assert torch.allclose(after, before, atol=1e-6) + + +def test_regular_lora_can_still_train_an_fft_allocated_layer(): + from twinkle.model.multi_lora import MultiLora + + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(TinyDecoder(num_layers=1), target_modules='all-linear') + hybrid = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + config = LoraConfig(r=2, lora_alpha=4, target_modules=['q_proj'], init_lora_weights=False) + manager.acquire_lora('regular', config) + tenant = manager.find_lora_by_tenant('regular') + q_proj = model.get_submodule(hybrid.allocated_to_layer_name['layers.0.self_attn.q_proj']) + inputs = torch.randn(2, 8) + + with manager.adapter('regular'): + before = model(inputs).detach() + with torch.no_grad(): + q_proj.lora_A[tenant.adapter_name].weight[:2].fill_(0.5) + q_proj.lora_B[tenant.adapter_name].weight[:, :2].fill_(0.5) + with manager.adapter('regular'): + after = model(inputs).detach() + + assert not torch.allclose(after, before) + assert hybrid._get_fft_wrapper('layers.0.self_attn.q_proj').active_adapters == [] + + +def _make_multi_tenant_hybrid(): + from twinkle.model.multi_lora import MultiLora + + torch.manual_seed(42) + base = TinyDecoder(num_layers=1, dim=8) + manager = MultiLora(max_loras=2, max_r=4) + model = manager.patch(base, target_modules='all-linear') + spectral = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + manager.save_initial_weights() + hybrid_config = LoraConfig( + r=2, + lora_alpha=4, + target_modules=['layers.0.mlp.down_proj'], + modules_to_save=['layers.0.self_attn.q_proj'], + init_lora_weights=False, + ) + regular_config = LoraConfig( + r=2, + lora_alpha=4, + target_modules=['layers.0.self_attn.v_proj'], + init_lora_weights=False, + ) + _register_hybrid(manager, spectral, 'hybrid', hybrid_config) + manager.acquire_lora('regular', regular_config) + return base, model, manager, spectral + + +def test_multi_tenant_hybrid_isolation_and_non_destructive_merge(): + base, model, manager, spectral = _make_multi_tenant_hybrid() + inputs = torch.randn(2, 8) + with manager.adapter('regular'): + regular_before = model(inputs).detach().clone() + + hybrid = manager.find_lora_by_tenant('hybrid') + modules = dict(model.named_modules()) + down = next( + layer for name, layer in modules.items() + if name.endswith('layers.0.mlp.down_proj') + ) + q_proj = _fft_slot_module(spectral, 'layers.0.self_attn.q_proj', hybrid.index) + with torch.no_grad(): + down.lora_A[hybrid.adapter_name].weight[:2].fill_(0.15) + down.lora_B[hybrid.adapter_name].weight[:, :2].fill_(0.10) + q_proj.weight.add_(0.05) + + before_export = {name: value.detach().clone() for name, value in model.named_parameters()} + with _adapter(manager, spectral, 'hybrid'): + hybrid_output = model(inputs).detach() + full_state = {name: value.detach().cpu().clone() for name, value in model.named_parameters()} + merged_state = spectral.build_merged_state_dict('hybrid', full_state) + + deployed = TinyDecoder(num_layers=1, dim=8) + deployed.load_state_dict(merged_state) + assert torch.allclose(deployed(inputs), hybrid_output, atol=1e-5) + assert all(torch.equal(before_export[name], value) for name, value in model.named_parameters()) + + with manager.adapter('regular'): + assert torch.allclose(model(inputs), regular_before, atol=1e-6) + + +def test_disable_lora_disables_hybrid_fft_slot_too(): + _, model, manager, spectral = _make_multi_tenant_hybrid() + inputs = torch.randn(2, 8) + hybrid = manager.find_lora_by_tenant('hybrid') + fft_layer = _fft_slot_module(spectral, 'layers.0.self_attn.q_proj', hybrid.index) + with _adapter(manager, spectral, 'hybrid', disable_lora=True): + base = model(inputs).detach() + with torch.no_grad(): + fft_layer.weight.add_(1.0) + + with _adapter(manager, spectral, 'hybrid', disable_lora=True): + disabled = model(inputs).detach() + + assert torch.allclose(disabled, base, atol=1e-6) + + +def test_hybrid_training_state_round_trip_and_release_resets_fft_slot(): + _, model, manager, spectral = _make_multi_tenant_hybrid() + hybrid = manager.find_lora_by_tenant('hybrid') + fft_layer = _fft_slot_module(spectral, 'layers.0.self_attn.q_proj', hybrid.index) + with torch.no_grad(): + fft_layer.weight.add_(0.25) + full_state = {name: value.detach().cpu().clone() for name, value in model.named_parameters()} + saved = spectral.build_training_state_dict('hybrid', full_state) + + with torch.no_grad(): + fft_layer.weight.zero_() + manager.set_state_dict('hybrid', saved) + spectral.set_fft_state_dict('hybrid', saved) + assert torch.equal( + fft_layer.weight, + saved['base_model.model.layers.0.self_attn.q_proj.weight'], + ) + + spectral.reset_adapter_slot('hybrid') + spectral.unregister_adapter('hybrid') + manager.release_lora('hybrid') + wrapper = spectral._get_fft_wrapper('layers.0.self_attn.q_proj') + assert torch.equal(fft_layer.weight, wrapper.original_module.weight) + + +def test_fft_state_traversal_includes_buffers(): + from twinkle.model.multi_lora import MultiLora + + base = TinyDecoder(num_layers=1) + q_proj = base.layers[0].self_attn.q_proj + q_proj.register_buffer('calibration', torch.tensor([1.0, 2.0])) + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(base, target_modules='all-linear') + fft_slots = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) + config = LoraConfig(r=2, lora_alpha=4, target_modules=['down_proj']) + _register_hybrid(manager, fft_slots, 'hybrid', config) + fft_module = _fft_slot_module(fft_slots, 'layers.0.self_attn.q_proj') + state_key = 'base_model.model.layers.0.self_attn.q_proj.calibration' + + fft_module.calibration.fill_(3.0) + saved = fft_slots.get_fft_state_dict('hybrid') + assert torch.equal(saved[state_key], torch.tensor([3.0, 3.0])) + + fft_module.calibration.zero_() + fft_slots.set_fft_state_dict('hybrid', saved) + assert torch.equal(fft_module.calibration, torch.tensor([3.0, 3.0])) + + fft_slots.reset_adapter_slot('hybrid') + assert torch.equal(fft_module.calibration, torch.tensor([1.0, 2.0])) + + +def test_multi_tenant_optimizer_parameters_and_learning_rates_are_isolated(): + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + _, model, manager, spectral = _make_multi_tenant_hybrid() + wrapper = object.__new__(SpectralHybridTransformersModel) + wrapper.multi_adapter = manager + wrapper.fft_slots = spectral + wrapper.strategy = type('Strategy', (), {'unwrap_model': lambda _self, inner: inner})() + wrapper.__dict__['model'] = model + wrapper.default_lr_lora = 2.5e-5 + wrapper.default_lr_fft = 1e-6 + + regular = wrapper._get_trainable_parameters('regular') + hybrid = wrapper._get_trainable_parameters('hybrid') + groups = wrapper._create_param_group('hybrid', weight_decay=0.0) + + assert regular + assert all('.lora_' in name and '.lora_1.' in name for name in regular) + assert all('v_proj' in name for name in regular) + assert any('down_proj' in name and '.lora_0.' in name for name in hybrid) + assert any('.modules_to_save.fft_0.' in name for name in hybrid) + assert {group['lr'] for group in groups} == {2.5e-5, 1e-6} + assert {id(param) for group in groups for param in group['params']} == {id(param) for param in hybrid.values()} + + +def test_regular_lora_rejects_targets_outside_preallocation(): + _, _, manager, _ = _make_multi_tenant_hybrid() + + with pytest.raises(ValueError, match='outside the preallocated range'): + manager.validate_tenant_target_modules(['does_not_exist']) + + +def test_merged_only_hybrid_checkpoint_cannot_resume(tmp_path): + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + wrapper = object.__new__(SpectralHybridTransformersModel) + with pytest.raises(ValueError, match='merged-only checkpoint'): + wrapper._resume_spectral_hybrid(str(tmp_path), 'hybrid', False) + + +def test_hybrid_save_optimizer_requires_configured_optimizer(tmp_path): + from types import SimpleNamespace + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + _, model, manager, spectral = _make_multi_tenant_hybrid() + wrapper = object.__new__(SpectralHybridTransformersModel) + wrapper.multi_adapter = manager + wrapper.fft_slots = spectral + wrapper.__dict__['model'] = model + wrapper.optimizer_group = { + 'hybrid': SimpleNamespace( + cur_step=0, + gradient_accumulation_steps=1, + optimizer=None, + train_status=SimpleNamespace(loss_value=None, num_tokens=0), + ) + } + + with pytest.raises(ValueError, match='optimizer must be configured'): + wrapper._validate_hybrid_training_checkpoint_boundary('hybrid') + + +def test_hybrid_save_writes_deployable_model_and_lossless_training_state(tmp_path): + from types import SimpleNamespace + from transformers import PretrainedConfig + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + _, model, manager, spectral = _make_multi_tenant_hybrid() + hybrid = manager.find_lora_by_tenant('hybrid') + fft_layer = _fft_slot_module(spectral, 'layers.0.self_attn.q_proj', hybrid.index) + with torch.no_grad(): + fft_layer.weight.add_(0.2) + + class Strategy: + + @staticmethod + def unwrap_model(inner): + return inner + + @staticmethod + def get_full_state_dict(inner): + return {name: value.detach().cpu().clone() for name, value in inner.named_parameters()} + + @staticmethod + def needs_wrapped_optimizer_state(): + return False + + @staticmethod + def save_optimizer_checkpoint(_model, optimizer, output_path): + torch.save(optimizer.state_dict(), output_path) + + @staticmethod + def load_optimizer_checkpoint(_model, optimizer, input_path): + optimizer.load_state_dict(torch.load(input_path, map_location='cpu', weights_only=False)) + + trainable = {} + for name, parameter in model.named_parameters(): + if '.lora_0.' in name or '.modules_to_save.fft_0.' in name: + trainable[name] = parameter + optimizer = torch.optim.AdamW([{ + 'params': list(trainable.values()), + 'param_names': list(trainable), + 'lr': 1e-3, + }]) + + group = SimpleNamespace( + cur_step=3, + gradient_accumulation_steps=2, + adapter_config=hybrid.tenant_config, + optimizer=optimizer, + lr_scheduler=None, + scaler=None, + train_status=SimpleNamespace(loss_value=None, num_tokens=0), + do_grad_sync=lambda: True, + ) + wrapper = object.__new__(SpectralHybridTransformersModel) + wrapper.multi_adapter = manager + wrapper.fft_slots = spectral + wrapper.strategy = Strategy() + wrapper.__dict__['model'] = model + wrapper.optimizer_group = {'hybrid': group} + wrapper.hf_config = PretrainedConfig() + object.__setattr__(wrapper, '_save_tokenizer', lambda *_args, **_kwargs: None) + + checkpoint = wrapper._save_spectral_hybrid( + 'checkpoint-3', + str(tmp_path), + 1, + 'hybrid', + save_optimizer=True, + consumed_train_samples=17, + ) + + assert (tmp_path / 'checkpoint-3' / 'config.json').is_file() + assert (tmp_path / 'checkpoint-3' / 'model.safetensors').is_file() + assert not (tmp_path / 'checkpoint-3' / 'adapter_config.json').exists() + training_dir = tmp_path / 'checkpoint-3' / 'twinkle_training_state' + assert (training_dir / 'adapter_config.json').is_file() + assert '"twinkle_adapter_mode": "hybrid"' in (training_dir / 'adapter_config.json').read_text() + assert (training_dir / 'adapter_model.safetensors').is_file() + assert (training_dir / 'optimizer.pt').is_file() + assert (training_dir / 'trainer_state.json').is_file() + assert (training_dir / 'rng_state_rank_0.pt').is_file() + + expected_fft = fft_layer.weight.detach().clone() + with torch.no_grad(): + fft_layer.weight.zero_() + progress = wrapper._resume_spectral_hybrid(checkpoint, 'hybrid', resume_only_model=True) + assert torch.equal(fft_layer.weight, expected_fft) + assert progress == { + 'cur_step': 3, + 'consumed_train_samples': 17, + 'gradient_accumulation_steps': 2, + } + + hybrid.tenant_config.lora_alpha += 1 + unchanged = fft_layer.weight.detach().clone() + with pytest.raises(ValueError, match='config does not match checkpoint'): + wrapper._resume_spectral_hybrid(checkpoint, 'hybrid', resume_only_model=True) + assert torch.equal(fft_layer.weight, unchanged) + + hybrid.tenant_config.lora_alpha -= 1 + wrapper._save_spectral_hybrid('checkpoint-3', str(tmp_path), 1, 'hybrid', save_optimizer=False) + assert not training_dir.exists() + + +def test_hybrid_optimizer_resume_matches_uninterrupted_training(tmp_path): + from types import SimpleNamespace + from transformers import PretrainedConfig + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + class Strategy: + + @staticmethod + def unwrap_model(inner): + return inner + + @staticmethod + def get_full_state_dict(inner): + return {name: value.detach().cpu().clone() for name, value in inner.state_dict().items()} + + @staticmethod + def needs_wrapped_optimizer_state(): + return False + + @staticmethod + def save_optimizer_checkpoint(_model, optimizer, output_path): + torch.save(optimizer.state_dict(), output_path) + + @staticmethod + def load_optimizer_checkpoint(_model, optimizer, input_path): + optimizer.load_state_dict(torch.load(input_path, map_location='cpu', weights_only=False)) + + def build_wrapper(): + _, model, manager, spectral = _make_multi_tenant_hybrid() + wrapper = object.__new__(SpectralHybridTransformersModel) + wrapper.multi_adapter = manager + wrapper.fft_slots = spectral + wrapper.strategy = Strategy() + wrapper.__dict__['model'] = model + wrapper.default_lr_lora = 1e-3 + wrapper.default_lr_fft = 2e-4 + params = wrapper._create_param_group('hybrid', weight_decay=0.0) + optimizer = torch.optim.AdamW(params) + tenant = manager.find_lora_by_tenant('hybrid') + group = SimpleNamespace( + cur_step=0, + gradient_accumulation_steps=1, + adapter_config=tenant.tenant_config, + optimizer=optimizer, + lr_scheduler=None, + scaler=None, + train_status=SimpleNamespace(loss_value=None, num_tokens=0), + do_grad_sync=lambda: True, + ) + wrapper.optimizer_group = {'hybrid': group} + wrapper.hf_config = PretrainedConfig() + object.__setattr__(wrapper, '_save_tokenizer', lambda *_args, **_kwargs: None) + wrapper._model_wrapped = False + return wrapper + + def train_step(wrapper, inputs): + group = wrapper.optimizer_group['hybrid'] + with wrapper._adapter_context('hybrid'): + loss = wrapper.model(inputs).float().square().mean() + loss.backward() + group.optimizer.step() + group.optimizer.zero_grad(set_to_none=True) + group.cur_step += 1 + + uninterrupted = build_wrapper() + train_step(uninterrupted, torch.full((2, 8), 0.25)) + checkpoint = uninterrupted._save_spectral_hybrid( + 'resume', str(tmp_path), 1, 'hybrid', save_optimizer=True, consumed_train_samples=2) + + resumed = build_wrapper() + progress = resumed._resume_spectral_hybrid(checkpoint, 'hybrid', resume_only_model=False) + assert progress['cur_step'] == 1 + + next_inputs = torch.full((2, 8), -0.4) + train_step(uninterrupted, next_inputs) + train_step(resumed, next_inputs) + + uninterrupted_state = uninterrupted.multi_adapter.get_state_dict('hybrid') + resumed_state = resumed.multi_adapter.get_state_dict('hybrid') + assert set(uninterrupted_state) == set(resumed_state) + for name in uninterrupted_state: + assert torch.equal(uninterrupted_state[name], resumed_state[name]), name + + uninterrupted_optimizer = uninterrupted.optimizer_group['hybrid'].optimizer.state_dict() + resumed_optimizer = resumed.optimizer_group['hybrid'].optimizer.state_dict() + assert uninterrupted_optimizer['param_groups'] == resumed_optimizer['param_groups'] + for parameter_id, state in uninterrupted_optimizer['state'].items(): + for state_name, value in state.items(): + restored = resumed_optimizer['state'][parameter_id][state_name] + if isinstance(value, torch.Tensor): + assert torch.equal(value, restored) + else: + assert value == restored + + +def test_hybrid_checkpoint_reloads_with_auto_model(tmp_path): + from types import SimpleNamespace + from transformers import AutoModelForCausalLM, LlamaConfig, LlamaForCausalLM + from twinkle.model.multi_lora import MultiLora + from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel + + config = LlamaConfig( + vocab_size=32, + hidden_size=8, + intermediate_size=16, + num_hidden_layers=1, + num_attention_heads=1, + num_key_value_heads=1, + max_position_embeddings=32, + tie_word_embeddings=True, + ) + manager = MultiLora(max_loras=1, max_r=4) + model = manager.patch(LlamaForCausalLM(config), target_modules='all-linear') + spectral = _install_hybrid(manager, model, ['model.layers.0.self_attn.q_proj']) + manager.save_initial_weights() + tenant_config = LoraConfig( + r=2, + lora_alpha=4, + target_modules=['model.layers.0.mlp.down_proj'], + modules_to_save=['model.layers.0.self_attn.q_proj'], + init_lora_weights=False, + ) + _register_hybrid(manager, spectral, 'hybrid', tenant_config) + tenant = manager.find_lora_by_tenant('hybrid') + modules = dict(model.named_modules()) + with torch.no_grad(): + _fft_slot_module(spectral, 'model.layers.0.self_attn.q_proj').weight.add_(0.02) + down = next( + layer for name, layer in modules.items() + if name.endswith('model.layers.0.mlp.down_proj') + ) + down.lora_A[tenant.adapter_name].weight[:2].fill_(0.1) + down.lora_B[tenant.adapter_name].weight[:, :2].fill_(0.1) + model.eval() + input_ids = torch.tensor([[1, 2, 3, 4]]) + with _adapter(manager, spectral, 'hybrid'), torch.no_grad(): + expected = model(input_ids=input_ids).logits + + class Strategy: + + @staticmethod + def unwrap_model(inner): + return inner + + @staticmethod + def get_full_state_dict(inner): + return {name: value.detach().cpu().clone() for name, value in inner.named_parameters()} + + wrapper = object.__new__(SpectralHybridTransformersModel) + wrapper.multi_adapter = manager + wrapper.fft_slots = spectral + wrapper.strategy = Strategy() + wrapper.__dict__['model'] = model + wrapper.optimizer_group = {'hybrid': SimpleNamespace(cur_step=0)} + wrapper.hf_config = config + object.__setattr__(wrapper, '_save_tokenizer', lambda *_args, **_kwargs: None) + checkpoint = wrapper._save_spectral_hybrid( + 'deploy', str(tmp_path), 1, 'hybrid', save_optimizer=False, max_shard_size='1KB') + + assert (tmp_path / 'deploy' / 'model.safetensors.index.json').is_file() + assert len(list((tmp_path / 'deploy').glob('model-*.safetensors'))) > 1 + deployed = AutoModelForCausalLM.from_pretrained(checkpoint).eval() + with torch.no_grad(): + actual = deployed(input_ids=input_ids).logits + assert torch.allclose(actual, expected, atol=1e-5) + + wrapper._save_spectral_hybrid( + 'deploy', str(tmp_path), 1, 'hybrid', save_optimizer=False, max_shard_size='5GB') + assert not (tmp_path / 'deploy' / 'model.safetensors.index.json').exists() + assert (tmp_path / 'deploy' / 'model.safetensors').is_file() + assert not list((tmp_path / 'deploy').glob('model-*.safetensors')) + reloaded = AutoModelForCausalLM.from_pretrained(checkpoint).eval() + with torch.no_grad(): + assert torch.allclose(reloaded(input_ids=input_ids).logits, expected, atol=1e-5) From 34532b997efa6b4f542bdc39bd73579d5a216616 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 14 Aug 2026 14:15:16 +0800 Subject: [PATCH 5/9] wip --- .../model/transformers/hybrid/fft_slots.py | 23 +++++++++++++++++ .../model/transformers/strategy/accelerate.py | 25 ------------------- 2 files changed, 23 insertions(+), 25 deletions(-) diff --git a/src/twinkle/model/transformers/hybrid/fft_slots.py b/src/twinkle/model/transformers/hybrid/fft_slots.py index cb77fd7ad..11c6dea1d 100644 --- a/src/twinkle/model/transformers/hybrid/fft_slots.py +++ b/src/twinkle/model/transformers/hybrid/fft_slots.py @@ -18,6 +18,7 @@ class HybridFftSlots: """ def __init__(self, multi_lora: MultiLora, s_fft: List[str]) -> None: + """Bind the shared LoRA manager to its server-owned FFT allocation.""" if not s_fft: raise ValueError('Hybrid requires at least one S_FFT module.') if len(set(s_fft)) != len(s_fft): @@ -33,10 +34,12 @@ def __init__(self, multi_lora: MultiLora, s_fft: List[str]) -> None: @property def module(self): + """Return the PEFT model that owns the FFT wrappers.""" return self.multi_lora.module @staticmethod def _canonical_module_name(name: str) -> str: + """Remove the module name prefix introduced by PEFT.""" prefix = 'base_model.model.' return name[len(prefix):] if name.startswith(prefix) else name @@ -60,6 +63,7 @@ def install_fft_slots(self) -> None: raise ValueError( f'Hybrid S_FFT aliases resolve to the same layer {layer_name!r}.') resolved_layer_names.add(layer_name) + if isinstance(layer, LoraLayer): original_module = layer.base_layer wrapper_name = f'{layer_name}.base_layer' @@ -84,32 +88,41 @@ def install_fft_slots(self) -> None: self.allocated_to_wrapper_name[allocated_name] = wrapper_name def is_hybrid(self, adapter_name: str) -> bool: + """Return whether an adapter owns an FFT slot.""" return adapter_name in self.hybrid_adapters def register_adapter(self, adapter_name: str) -> None: + """Mark an existing LoRA tenant as a Hybrid tenant.""" self.hybrid_adapters.add(adapter_name) def unregister_adapter(self, adapter_name: str) -> None: + """Remove Hybrid ownership before releasing a tenant slot.""" self.hybrid_adapters.discard(adapter_name) def _tenant(self, adapter_name: str): + """Look up the LoRA tenant that determines the FFT slot index.""" return self.multi_lora.find_lora_by_tenant(adapter_name) def _fft_adapter_name(self, adapter_name: str) -> str: + """Map a tenant's LoRA slot index to its PEFT FFT adapter name.""" return f'fft_{self._tenant(adapter_name).index}' def _get_fft_wrapper(self, allocated_name: str) -> ModulesToSaveWrapper: + """Resolve one allocated module to its installed PEFT wrapper.""" return self.module.get_submodule(self.allocated_to_wrapper_name[allocated_name]) def _iter_fft_wrappers(self): + """Return all wrappers in stable allocation order.""" return [self._get_fft_wrapper(name) for name in self.s_fft] def activate_fft_slot(self, adapter_name: str) -> None: + """Activate this tenant's FFT copies, or disable FFT for regular LoRA.""" fft_adapter_name = self._fft_adapter_name(adapter_name) if self.is_hybrid(adapter_name) else None for wrapper in self._iter_fft_wrappers(): wrapper.set_adapter(fft_adapter_name if fft_adapter_name is not None else []) def deactivate_fft_slots(self) -> None: + """Disable every FFT wrapper after a model operation.""" for wrapper in self._iter_fft_wrappers(): wrapper.set_adapter([]) @@ -127,6 +140,7 @@ def resolve_lora_targets(self, target_modules) -> List[str]: return sorted(name for name in selected if name not in fft_layers) def _tenant_lora_layer_names(self, adapter_name: str) -> List[str]: + """Return this Hybrid tenant's LoRA layers, excluding its FFT layers.""" tenant = self._tenant(adapter_name) fft_layers = set(self.allocated_to_layer_name.values()) return [ @@ -137,12 +151,14 @@ def _tenant_lora_layer_names(self, adapter_name: str) -> List[str]: @staticmethod def _iter_module_tensors(module: nn.Module): + """Yield a module's parameters and buffers with their persistence kind.""" for name, parameter in module.named_parameters(): yield name, parameter, True for name, buffer in module.named_buffers(): yield name, buffer, False def _iter_fft_slot_tensors(self, adapter_name: str): + """Yield all tensors belonging to one tenant's FFT module copies.""" fft_adapter_name = self._fft_adapter_name(adapter_name) for allocated_name in self.s_fft: wrapper_name = self.allocated_to_wrapper_name[allocated_name] @@ -151,6 +167,7 @@ def _iter_fft_slot_tensors(self, adapter_name: str): yield allocated_name, wrapper_name, fft_adapter_name, tensor_name, tensor, is_parameter def reset_adapter_slot(self, adapter_name: str) -> None: + """Restore this tenant's FFT copies from the frozen original modules.""" for allocated_name, _, _, tensor_name, target, _ in self._iter_fft_slot_tensors(adapter_name): wrapper = self._get_fft_wrapper(allocated_name) original_tensors = { @@ -162,9 +179,11 @@ def reset_adapter_slot(self, adapter_name: str) -> None: @staticmethod def _checkpoint_key(allocated_name: str, parameter_name: str) -> str: + """Build the plain Transformers checkpoint key for an FFT tensor.""" return f'base_model.model.{allocated_name}.{parameter_name}' def get_fft_state_dict(self, adapter_name: str) -> Dict[str, torch.Tensor]: + """Return an independent snapshot of one tenant's FFT state.""" if not self.is_hybrid(adapter_name): return {} state = {} @@ -174,6 +193,7 @@ def get_fft_state_dict(self, adapter_name: str) -> Dict[str, torch.Tensor]: return state def set_fft_state_dict(self, adapter_name: str, state_dict: Dict[str, torch.Tensor]) -> None: + """Restore one tenant's FFT copies from lossless training state.""" if not self.is_hybrid(adapter_name): return for allocated_name, _, _, tensor_name, tensor, _ in self._iter_fft_slot_tensors(adapter_name): @@ -183,6 +203,7 @@ def set_fft_state_dict(self, adapter_name: str, state_dict: Dict[str, torch.Tens self.multi_lora._write_param_tensor(tensor, state_dict[key]) def named_fft_parameters(self, adapter_name: str): + """Return trainable FFT parameters with their PEFT state-dict names.""" if not self.is_hybrid(adapter_name): return [] result = [] @@ -194,6 +215,7 @@ def named_fft_parameters(self, adapter_name: str): @staticmethod def _normalize_base_state_key(name: str) -> str: + """Convert a PEFT base-module key to its plain Transformers key.""" prefix = 'base_model.model.' if name.startswith(prefix): name = name[len(prefix):] @@ -246,6 +268,7 @@ def iter_merged_state_dict(self, adapter_name: str, full_state_dict: Dict[str, t yield output_name, replacements.get(name, value).detach().cpu() def build_merged_state_dict(self, adapter_name: str, full_state_dict: Dict[str, torch.Tensor]): + """Materialize the merged state dict for deployment-oriented saving.""" return dict(self.iter_merged_state_dict(adapter_name, full_state_dict)) def build_training_state_dict(self, adapter_name: str, full_state_dict: Dict[str, torch.Tensor]): diff --git a/src/twinkle/model/transformers/strategy/accelerate.py b/src/twinkle/model/transformers/strategy/accelerate.py index a8e77669b..e6314b531 100644 --- a/src/twinkle/model/transformers/strategy/accelerate.py +++ b/src/twinkle/model/transformers/strategy/accelerate.py @@ -147,20 +147,6 @@ def _get_fsdp_plugin(self): state = self.accelerator.state return state.fsdp_plugin if hasattr(state, 'fsdp_plugin') else None - def _prepare_fsdp2_sd_options(self): - fsdp_plugin = self._get_fsdp_plugin() - if fsdp_plugin is None or fsdp_plugin.fsdp_version != 2: - return None - - from torch.distributed.checkpoint.state_dict import StateDictOptions - from torch.distributed.fsdp.fully_sharded_data_parallel import StateDictType - - return StateDictOptions( - full_state_dict=fsdp_plugin.state_dict_type == StateDictType.FULL_STATE_DICT, - cpu_offload=getattr(fsdp_plugin.state_dict_config, 'offload_to_cpu', False), - broadcast_from_rank0=getattr(fsdp_plugin.state_dict_config, 'rank0_only', False), - ) - @staticmethod def _prepare_full_optimizer_state_dict_options(*, for_load: bool): from torch.distributed.checkpoint.state_dict import StateDictOptions @@ -215,17 +201,6 @@ def load_optimizer_checkpoint(self, model, optimizer, input_path: str): def get_full_state_dict(self, model) -> dict: """Collect full state dict.""" from twinkle.utils import torch_util - fsdp_plugin = self._get_fsdp_plugin() - if fsdp_plugin is not None and fsdp_plugin.fsdp_version == 2: - from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict - - # A merged checkpoint is a deployment artifact, so it must always - # materialize global, CPU-offloaded parameters on rank 0 regardless - # of the training plugin's normal sharded checkpoint preference. - return get_model_state_dict( - model, - options=StateDictOptions(full_state_dict=True, cpu_offload=True), - ) unwrapped = self.unwrap_model(model) state_dict = {} for name, value in unwrapped.state_dict().items(): From 6099238aa0b030c223112b91fff01e5ab248a6d3 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 14 Aug 2026 14:30:13 +0800 Subject: [PATCH 6/9] wip --- .../model/transformers/hybrid/fft_slots.py | 60 ++++++--------- .../model/transformers/strategy/accelerate.py | 4 +- .../transformers/strategy/native_fsdp.py | 61 ++++++--------- .../transformers/test_spectral_hybrid_lora.py | 74 ------------------- 4 files changed, 50 insertions(+), 149 deletions(-) diff --git a/src/twinkle/model/transformers/hybrid/fft_slots.py b/src/twinkle/model/transformers/hybrid/fft_slots.py index 11c6dea1d..6be157cb3 100644 --- a/src/twinkle/model/transformers/hybrid/fft_slots.py +++ b/src/twinkle/model/transformers/hybrid/fft_slots.py @@ -149,33 +149,22 @@ def _tenant_lora_layer_names(self, adapter_name: str) -> List[str]: and self.multi_lora.match_target_modules(name, tenant.tenant_config.target_modules) ] - @staticmethod - def _iter_module_tensors(module: nn.Module): - """Yield a module's parameters and buffers with their persistence kind.""" - for name, parameter in module.named_parameters(): - yield name, parameter, True - for name, buffer in module.named_buffers(): - yield name, buffer, False - - def _iter_fft_slot_tensors(self, adapter_name: str): - """Yield all tensors belonging to one tenant's FFT module copies.""" + def _iter_fft_slot_parameters(self, adapter_name: str): + """Yield all parameters belonging to one tenant's FFT module copies.""" fft_adapter_name = self._fft_adapter_name(adapter_name) for allocated_name in self.s_fft: wrapper_name = self.allocated_to_wrapper_name[allocated_name] slot_module = self._get_fft_wrapper(allocated_name).modules_to_save[fft_adapter_name] - for tensor_name, tensor, is_parameter in self._iter_module_tensors(slot_module): - yield allocated_name, wrapper_name, fft_adapter_name, tensor_name, tensor, is_parameter + for parameter_name, parameter in slot_module.named_parameters(): + yield allocated_name, wrapper_name, fft_adapter_name, parameter_name, parameter def reset_adapter_slot(self, adapter_name: str) -> None: """Restore this tenant's FFT copies from the frozen original modules.""" - for allocated_name, _, _, tensor_name, target, _ in self._iter_fft_slot_tensors(adapter_name): + for allocated_name, _, _, parameter_name, target in self._iter_fft_slot_parameters(adapter_name): wrapper = self._get_fft_wrapper(allocated_name) - original_tensors = { - name: tensor - for name, tensor, _ in self._iter_module_tensors(wrapper.original_module) - } + original_parameters = dict(wrapper.original_module.named_parameters()) self.multi_lora._write_param_tensor( - target, self.multi_lora._read_param_tensor(original_tensors[tensor_name])) + target, self.multi_lora._read_param_tensor(original_parameters[parameter_name])) @staticmethod def _checkpoint_key(allocated_name: str, parameter_name: str) -> str: @@ -187,31 +176,30 @@ def get_fft_state_dict(self, adapter_name: str) -> Dict[str, torch.Tensor]: if not self.is_hybrid(adapter_name): return {} state = {} - for allocated_name, _, _, tensor_name, tensor, _ in self._iter_fft_slot_tensors(adapter_name): - state[self._checkpoint_key(allocated_name, tensor_name)] = ( - self.multi_lora._read_param_tensor(tensor).detach().clone()) + for allocated_name, _, _, parameter_name, parameter in self._iter_fft_slot_parameters(adapter_name): + state[self._checkpoint_key(allocated_name, parameter_name)] = ( + self.multi_lora._read_param_tensor(parameter).detach().clone()) return state def set_fft_state_dict(self, adapter_name: str, state_dict: Dict[str, torch.Tensor]) -> None: """Restore one tenant's FFT copies from lossless training state.""" if not self.is_hybrid(adapter_name): return - for allocated_name, _, _, tensor_name, tensor, _ in self._iter_fft_slot_tensors(adapter_name): - key = self._checkpoint_key(allocated_name, tensor_name) + for allocated_name, _, _, parameter_name, parameter in self._iter_fft_slot_parameters(adapter_name): + key = self._checkpoint_key(allocated_name, parameter_name) if key not in state_dict: raise ValueError(f'Hybrid training state is missing {key!r}.') - self.multi_lora._write_param_tensor(tensor, state_dict[key]) + self.multi_lora._write_param_tensor(parameter, state_dict[key]) def named_fft_parameters(self, adapter_name: str): """Return trainable FFT parameters with their PEFT state-dict names.""" if not self.is_hybrid(adapter_name): return [] - result = [] - for _, wrapper_name, fft_adapter_name, tensor_name, tensor, is_parameter in self._iter_fft_slot_tensors( - adapter_name): - if is_parameter: - result.append((f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}', tensor)) - return result + return [ + (f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{parameter_name}', parameter) + for _, wrapper_name, fft_adapter_name, parameter_name, parameter + in self._iter_fft_slot_parameters(adapter_name) + ] @staticmethod def _normalize_base_state_key(name: str) -> str: @@ -248,10 +236,10 @@ def iter_merged_state_dict(self, adapter_name: str, full_state_dict: Dict[str, t delta = delta.transpose(0, 1) replacements[base_key] = base + delta.to(dtype=base.dtype) * scaling - for allocated_name, wrapper_name, fft_adapter_name, tensor_name, _, _ in self._iter_fft_slot_tensors( + for allocated_name, wrapper_name, fft_adapter_name, parameter_name, _ in self._iter_fft_slot_parameters( adapter_name): - base_key = f'{wrapper_name}.original_module.{tensor_name}' - fft_key = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}' + base_key = f'{wrapper_name}.original_module.{parameter_name}' + fft_key = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{parameter_name}' if base_key not in full_state_dict or fft_key not in full_state_dict: raise ValueError(f'Cannot export Hybrid FFT layer {allocated_name!r}.') replacements[base_key] = full_state_dict[fft_key] @@ -285,10 +273,10 @@ def build_training_state_dict(self, adapter_name: str, full_state_dict: Dict[str value = self.multi_lora._slice_rank_tensor( source, full_state_dict[source], tenant.tenant_config.r) state[source.replace(f'.{tenant.adapter_name}.', '.')] = value.detach().cpu() - for allocated_name, wrapper_name, fft_adapter_name, tensor_name, _, _ in self._iter_fft_slot_tensors( + for allocated_name, wrapper_name, fft_adapter_name, parameter_name, _ in self._iter_fft_slot_parameters( adapter_name): - source = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{tensor_name}' + source = f'{wrapper_name}.modules_to_save.{fft_adapter_name}.{parameter_name}' if source not in full_state_dict: raise ValueError(f'Hybrid training state is missing {source!r}.') - state[self._checkpoint_key(allocated_name, tensor_name)] = full_state_dict[source].detach().cpu() + state[self._checkpoint_key(allocated_name, parameter_name)] = full_state_dict[source].detach().cpu() return state diff --git a/src/twinkle/model/transformers/strategy/accelerate.py b/src/twinkle/model/transformers/strategy/accelerate.py index e6314b531..4db3aa60d 100644 --- a/src/twinkle/model/transformers/strategy/accelerate.py +++ b/src/twinkle/model/transformers/strategy/accelerate.py @@ -203,8 +203,8 @@ def get_full_state_dict(self, model) -> dict: from twinkle.utils import torch_util unwrapped = self.unwrap_model(model) state_dict = {} - for name, value in unwrapped.state_dict().items(): - local = torch_util.to_local_tensor(value) + for name, param in unwrapped.named_parameters(): + local = torch_util.to_local_tensor(param) state_dict[name] = local.cpu() del local return state_dict diff --git a/src/twinkle/model/transformers/strategy/native_fsdp.py b/src/twinkle/model/transformers/strategy/native_fsdp.py index cce92bd03..92dbb8d62 100644 --- a/src/twinkle/model/transformers/strategy/native_fsdp.py +++ b/src/twinkle/model/transformers/strategy/native_fsdp.py @@ -318,45 +318,32 @@ def get_full_state_dict(self, model) -> dict: the local expert shards across the EP group to reconstruct the full expert tensor (all num_experts on dim-0). """ - if self.device_mesh is not None: - from torch.distributed.checkpoint.state_dict import StateDictOptions, get_model_state_dict - - ep_mesh = self.ep_fsdp_device_mesh - ep_world_size = ep_mesh['ep'].size() if ep_mesh is not None else 1 - if ep_world_size <= 1: - # FSDP2 parameters must be gathered through the state-dict API; - # CPU offload makes the result rank0-only. - return get_model_state_dict( - model, - options=StateDictOptions(full_state_dict=True, cpu_offload=True), - ) - - # EP experts are independently sharded across the EP dimension, so - # every EP rank must retain its FSDP-gathered expert block until the - # second all-gather below has reconstructed the original tensor. - state_dict = get_model_state_dict( - model, - options=StateDictOptions(full_state_dict=True, cpu_offload=False), - ) - unwrapped = self.unwrap_model(model) - ep_expert_names = _detect_ep_expert_names(unwrapped) - ep_group = ep_mesh['ep'].get_group() - result = {} - for name, value in state_dict.items(): - if name in ep_expert_names: - local_full = value.contiguous().to(Platform.get_local_device()) - gathered = [torch.empty_like(local_full) for _ in range(ep_world_size)] - dist.all_gather(gathered, local_full, group=ep_group) - value = torch.cat(gathered, dim=_ep_expert_state_dict_gather_dim(name)) - if Platform.is_master(): - result[name] = value.cpu() - return result unwrapped = self.unwrap_model(model) state_dict = {} - for name, value in unwrapped.state_dict().items(): - local_full = torch_util.to_local_tensor(value) - state_dict[name] = local_full.cpu() - del local_full + + ep_fsdp_mesh = self.ep_fsdp_device_mesh + ep_group = None + ep_world_size = 1 + if ep_fsdp_mesh is not None: + ep_group = ep_fsdp_mesh['ep'].get_group() + ep_world_size = ep_fsdp_mesh['ep'].size() + + ep_expert_names = _detect_ep_expert_names(unwrapped) if ep_world_size > 1 else set() + + for name, param in unwrapped.named_parameters(): + local_full = torch_util.to_local_tensor(param) + + if name in ep_expert_names and ep_world_size > 1 and ep_group is not None: + local_full = local_full.contiguous().to(Platform.get_local_device()) + gathered = [torch.empty_like(local_full) for _ in range(ep_world_size)] + dist.all_gather(gathered, local_full, group=ep_group) + local_full = torch.cat(gathered, dim=_ep_expert_state_dict_gather_dim(name)) + state_dict[name] = local_full.cpu() + del gathered, local_full + else: + state_dict[name] = local_full.cpu() + del local_full + return state_dict def get_adapter_state_dict(self, model, adapter_name: str) -> dict: diff --git a/tests/transformers/test_spectral_hybrid_lora.py b/tests/transformers/test_spectral_hybrid_lora.py index 06edab19b..acf295e00 100644 --- a/tests/transformers/test_spectral_hybrid_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -328,54 +328,6 @@ def fake_set(_model, _optimizer, state, *, options): assert load_options.broadcast_from_rank0 is True -def test_native_fsdp_full_state_reconstructs_ep_experts(monkeypatch): - import torch.distributed.checkpoint.state_dict as state_dict_api - from twinkle import Platform - from twinkle.model.transformers.strategy import native_fsdp - from twinkle.model.transformers.strategy.native_fsdp import NativeFSDPStrategy - - class EpDimension: - - @staticmethod - def size(): - return 2 - - @staticmethod - def get_group(): - return 'ep-group' - - strategy = object.__new__(NativeFSDPStrategy) - strategy.device_mesh = object() - strategy.ep_fsdp_device_mesh = {'ep': EpDimension()} - strategy.unwrap_model = lambda model: model - captured_options = [] - - def get_model_state_dict(_model, *, options): - captured_options.append(options) - return { - 'experts.weight': torch.tensor([[1.0]]), - 'dense.weight': torch.tensor([[3.0]]), - } - - def all_gather(output, value, *, group): - assert group == 'ep-group' - output[0].copy_(value) - output[1].copy_(value + 1) - - monkeypatch.setattr(state_dict_api, 'get_model_state_dict', get_model_state_dict) - monkeypatch.setattr(native_fsdp, '_detect_ep_expert_names', lambda _model: {'experts.weight'}) - monkeypatch.setattr(native_fsdp.dist, 'all_gather', all_gather) - monkeypatch.setattr(Platform, 'get_local_device', lambda: 'cpu') - monkeypatch.setattr(Platform, 'is_master', lambda: True) - - state = strategy.get_full_state_dict(object()) - - assert captured_options[0].full_state_dict is True - assert captured_options[0].cpu_offload is False - assert torch.equal(state['experts.weight'], torch.tensor([[1.0], [2.0]])) - assert torch.equal(state['dense.weight'], torch.tensor([[3.0]])) - - def test_twinkle_checkpoint_normalization_round_trips_full_modules(tmp_path): from safetensors.torch import save_file from twinkle.model.transformers.strategy.accelerate import AccelerateStrategy @@ -764,32 +716,6 @@ def test_hybrid_training_state_round_trip_and_release_resets_fft_slot(): assert torch.equal(fft_layer.weight, wrapper.original_module.weight) -def test_fft_state_traversal_includes_buffers(): - from twinkle.model.multi_lora import MultiLora - - base = TinyDecoder(num_layers=1) - q_proj = base.layers[0].self_attn.q_proj - q_proj.register_buffer('calibration', torch.tensor([1.0, 2.0])) - manager = MultiLora(max_loras=1, max_r=4) - model = manager.patch(base, target_modules='all-linear') - fft_slots = _install_hybrid(manager, model, ['layers.0.self_attn.q_proj']) - config = LoraConfig(r=2, lora_alpha=4, target_modules=['down_proj']) - _register_hybrid(manager, fft_slots, 'hybrid', config) - fft_module = _fft_slot_module(fft_slots, 'layers.0.self_attn.q_proj') - state_key = 'base_model.model.layers.0.self_attn.q_proj.calibration' - - fft_module.calibration.fill_(3.0) - saved = fft_slots.get_fft_state_dict('hybrid') - assert torch.equal(saved[state_key], torch.tensor([3.0, 3.0])) - - fft_module.calibration.zero_() - fft_slots.set_fft_state_dict('hybrid', saved) - assert torch.equal(fft_module.calibration, torch.tensor([3.0, 3.0])) - - fft_slots.reset_adapter_slot('hybrid') - assert torch.equal(fft_module.calibration, torch.tensor([1.0, 2.0])) - - def test_multi_tenant_optimizer_parameters_and_learning_rates_are_isolated(): from twinkle.model.transformers.hybrid import SpectralHybridTransformersModel From 34e44a6cd8ec3c4c816830d7c6309bf91ad33a30 Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 14 Aug 2026 14:47:33 +0800 Subject: [PATCH 7/9] wip --- .../transformers/multi_lora_transformers.py | 19 ++----------------- .../test_multi_lora_target_parameters.py | 1 - .../transformers/test_spectral_hybrid_lora.py | 1 - 3 files changed, 2 insertions(+), 19 deletions(-) diff --git a/src/twinkle/model/transformers/multi_lora_transformers.py b/src/twinkle/model/transformers/multi_lora_transformers.py index 00e1b409b..701f1370f 100644 --- a/src/twinkle/model/transformers/multi_lora_transformers.py +++ b/src/twinkle/model/transformers/multi_lora_transformers.py @@ -360,20 +360,5 @@ def _get_trainable_parameters_example(self, adapter_name, model): return self.multi_adapter.get_trainable_parameters_example(adapter_name) def _get_trainable_parameters(self, adapter_name): - with self._adapter_context(adapter_name) as real_adapter_name: - tenant = self.multi_adapter.find_lora_by_tenant(adapter_name) - pattern = f'.{real_adapter_name}.' - params = {} - model = self.strategy.unwrap_model(self.model) - for name, parameter in model.named_parameters(): - if not parameter.requires_grad: - continue - if pattern in name and '.lora_' in name: - if self.multi_adapter.match_target_modules(name, tenant.tenant_config.target_modules): - params[name] = parameter - known_parameter_ids = {id(parameter) for parameter in params.values()} - for name, parameter in self.multi_adapter.target_parameter_manager.named_slot_parameters(adapter_name): - if id(parameter) not in known_parameter_ids: - params[name] = parameter - known_parameter_ids.add(id(parameter)) - return params + with self.multi_adapter.adapter(adapter_name) as real_adapter_name: + return super()._get_trainable_parameters(real_adapter_name) diff --git a/tests/model/test_multi_lora_target_parameters.py b/tests/model/test_multi_lora_target_parameters.py index e6a68e79a..0d2e5ffaa 100644 --- a/tests/model/test_multi_lora_target_parameters.py +++ b/tests/model/test_multi_lora_target_parameters.py @@ -278,7 +278,6 @@ def test_multilora_transformers_optimizer_includes_target_parameter_slots(): expected = dict(manager.target_parameter_manager.named_slot_parameters("adapter_a")) assert expected - assert set(selected) == set(expected) assert {id(value) for value in selected.values()} == {id(value) for value in expected.values()} diff --git a/tests/transformers/test_spectral_hybrid_lora.py b/tests/transformers/test_spectral_hybrid_lora.py index acf295e00..38965c3c8 100644 --- a/tests/transformers/test_spectral_hybrid_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -734,7 +734,6 @@ def test_multi_tenant_optimizer_parameters_and_learning_rates_are_isolated(): assert regular assert all('.lora_' in name and '.lora_1.' in name for name in regular) - assert all('v_proj' in name for name in regular) assert any('down_proj' in name and '.lora_0.' in name for name in hybrid) assert any('.modules_to_save.fft_0.' in name for name in hybrid) assert {group['lr'] for group in groups} == {2.5e-5, 1e-6} From 3d99079986b0c5e97296b0e758fa4399d4f2a90b Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 14 Aug 2026 14:50:46 +0800 Subject: [PATCH 8/9] wip --- src/twinkle/model/multi_lora_target_parameters.py | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) diff --git a/src/twinkle/model/multi_lora_target_parameters.py b/src/twinkle/model/multi_lora_target_parameters.py index 913b4d21e..0eb33eea7 100644 --- a/src/twinkle/model/multi_lora_target_parameters.py +++ b/src/twinkle/model/multi_lora_target_parameters.py @@ -364,25 +364,19 @@ def parameters_for_tenant(self, tenant_adapter_name: str) -> list[nn.Parameter]: return parameters def named_slot_parameters(self, tenant_adapter_name: str) -> Iterator[tuple[str, nn.Parameter]]: - slot_name = self.tenant_to_slot.get(tenant_adapter_name) - if slot_name is None: - return + slot_name = self.tenant_to_slot[tenant_adapter_name] for wrapper in self.wrappers: yield from wrapper.named_slot_parameters(slot_name) def get_state_dict(self, tenant_adapter_name: str) -> dict[str, torch.Tensor]: - slot_name = self.tenant_to_slot.get(tenant_adapter_name) - if slot_name is None: - return {} + slot_name = self.tenant_to_slot[tenant_adapter_name] state_dict = {} for wrapper in self.wrappers: state_dict.update(wrapper.get_state_dict(slot_name)) return state_dict def set_state_dict(self, tenant_adapter_name: str, state_dict: dict[str, torch.Tensor]) -> set[str]: - slot_name = self.tenant_to_slot.get(tenant_adapter_name) - if slot_name is None: - return set() + slot_name = self.tenant_to_slot[tenant_adapter_name] consumed_keys = set() for wrapper in self.wrappers: consumed_keys.update(wrapper.set_state_dict(slot_name, state_dict)) From 420ebbd816e82a1672fb1e71a4f9e6b2e10c717e Mon Sep 17 00:00:00 2001 From: weikaiwen <34648228+kevssim@users.noreply.github.com> Date: Fri, 14 Aug 2026 15:51:11 +0800 Subject: [PATCH 9/9] wip --- cookbook/transformers/spectral_hybrid_lora.py | 12 ++-- src/twinkle/model/base.py | 62 +++++++++---------- .../transformers/multi_lora_transformers.py | 4 -- src/twinkle/server/config/application_spec.py | 2 +- .../transformers/test_spectral_hybrid_lora.py | 21 ------- 5 files changed, 36 insertions(+), 65 deletions(-) diff --git a/cookbook/transformers/spectral_hybrid_lora.py b/cookbook/transformers/spectral_hybrid_lora.py index 87e79330f..3ba97f2eb 100644 --- a/cookbook/transformers/spectral_hybrid_lora.py +++ b/cookbook/transformers/spectral_hybrid_lora.py @@ -2,7 +2,6 @@ import os from pathlib import Path -import torch.distributed as dist from peft import LoraConfig import twinkle @@ -11,7 +10,6 @@ from twinkle.dataloader import DataLoader from twinkle.dataset import Dataset, DatasetMeta from twinkle.model import TransformersModel -from twinkle.model.base import initialize_process_group from twinkle.model.transformers.spectral_hybrid_lora import ( CANDIDATE_TYPES, allocate_spectral_modules, @@ -166,11 +164,11 @@ def resolve_spectral_config() -> LoraConfig: else: logger.info(f'No spectral config supplied; computing allocation and writing {config_path}') - initialize_process_group() - if Platform.is_master(): - compute_allocation(config_path) - if dist.is_available() and dist.is_initialized(): - dist.barrier() + if Platform.get_world_size() > 1: + raise ValueError( + 'Spectral allocation must be generated before multi-process training. ' + 'Run this cookbook once with a single process, then reuse the generated config.') + compute_allocation(config_path) return load_spectral_config(config_path) diff --git a/src/twinkle/model/base.py b/src/twinkle/model/base.py index d9e7dcbf1..8ea00d696 100644 --- a/src/twinkle/model/base.py +++ b/src/twinkle/model/base.py @@ -18,37 +18,6 @@ from torch.optim.lr_scheduler import LRScheduler -def initialize_process_group(should_bind_device_id: Optional[Callable[[str], bool]] = None) -> None: - """Initialize Twinkle's default distributed process group when launched with multiple ranks.""" - import torch - import torch.distributed as dist - if dist.is_initialized() or Platform.get_world_size() <= 1: - return - - torch_util.set_device() - backend = Platform.device_backend() - if backend == 'hccl': - # Keep training-side HCCL sockets on a per-job port layout to avoid collisions. - from twinkle.utils.platforms import ensure_hccl_socket_env - master_port = int(os.environ.get('MASTER_PORT', '29500')) - ensure_hccl_socket_env(master_port) - init_kwargs = { - 'backend': backend, - 'init_method': 'env://', - 'rank': Platform.get_rank(), - 'world_size': Platform.get_world_size(), - } - bind_device = should_bind_device_id(backend) if should_bind_device_id else backend in ('nccl', 'hccl') - if bind_device: - init_kwargs['device_id'] = torch.device(Platform.get_local_device()) - dist.init_process_group(**init_kwargs) - if backend == 'hccl': - # A bound HCCL default group can leak its device into later Gloo metric groups. - default_pg = dist.distributed_c10d._get_default_group() - if getattr(default_pg, 'bound_device_id', None) is not None: - default_pg.bound_device_id = None - - class TwinkleModel(ABC): _checkpoint_engine = None @@ -177,4 +146,33 @@ def _should_bind_device_id_for_process_group(self, backend: str) -> bool: return backend in ('nccl', 'hccl') def _try_init_process_group(self): - initialize_process_group(self._should_bind_device_id_for_process_group) + import torch + import torch.distributed as dist + if not dist.is_initialized() and Platform.get_world_size() > 1: + torch_util.set_device() + backend = Platform.device_backend() + if backend == 'hccl': + # fix: In multi-job NPU runs, HCCL default ports may collide (bind/listen failures). + # fix: Inject deterministic per-job port ranges before PG init to reduce cross-job conflicts. + # Keep training-side HCCL sockets on a per-job port layout to + # avoid collisions with other jobs on the same host. + from twinkle.utils.platforms import ensure_hccl_socket_env + master_port = int(os.environ.get('MASTER_PORT', '29500')) + ensure_hccl_socket_env(master_port) + init_kwargs = { + 'backend': backend, + 'init_method': 'env://', + 'rank': Platform.get_rank(), + 'world_size': Platform.get_world_size(), + } + if self._should_bind_device_id_for_process_group(backend): + init_kwargs['device_id'] = torch.device(Platform.get_local_device()) + dist.init_process_group(**init_kwargs) + if backend == 'hccl': + default_pg = dist.distributed_c10d._get_default_group() + if getattr(default_pg, 'bound_device_id', None) is not None: + # If the default HCCL PG keeps a bound device id, PyTorch may + # propagate that binding into later Gloo subgroup creation. That + # breaks the metrics/object-gather path on NPU, so clear it + # before Megatron creates its Gloo DP groups. + default_pg.bound_device_id = None diff --git a/src/twinkle/model/transformers/multi_lora_transformers.py b/src/twinkle/model/transformers/multi_lora_transformers.py index 701f1370f..488e9a4d0 100644 --- a/src/twinkle/model/transformers/multi_lora_transformers.py +++ b/src/twinkle/model/transformers/multi_lora_transformers.py @@ -41,7 +41,6 @@ def __init__( max_r: int = 32, max_length: int = 8192, target_modules: Union[List[str], str] = 'all-linear', - preallocated_lora_modules: Optional[Union[List[str], str]] = None, **kwargs): os.environ['TOKENIZERS_PARALLELISM'] = 'true' self._try_init_process_group() @@ -84,9 +83,6 @@ def __init__( self.sp_strategy = None # Initialize expert parallel attributes (required by set_optimizer in TransformersModel) self.optimizer_group: Dict[str, OptimizerGroup] = {} - if preallocated_lora_modules is not None: - target_modules = preallocated_lora_modules - self.preallocated_lora_modules = target_modules self.multi_adapter = MultiLora(max_loras=max_loras, max_r=max_r, max_length=max_length) self.model.gradient_checkpointing_enable() self.model = self.multi_adapter.patch(self.model, target_modules=target_modules, lora_config=self.lora_config) diff --git a/src/twinkle/server/config/application_spec.py b/src/twinkle/server/config/application_spec.py index 20db9dd84..9486c7656 100644 --- a/src/twinkle/server/config/application_spec.py +++ b/src/twinkle/server/config/application_spec.py @@ -75,7 +75,7 @@ class ModelArgs(_ArgsBase): max_loras: int = 5 max_r: int = Field(default=32, gt=0) max_length: int | None = None - preallocated_lora_modules: str | list[str] = 'all-linear' + target_modules: str | list[str] = 'all-linear' hybrid: SpectralHybridArgs | None = None @model_validator(mode='after') diff --git a/tests/transformers/test_spectral_hybrid_lora.py b/tests/transformers/test_spectral_hybrid_lora.py index 38965c3c8..a0e1af230 100644 --- a/tests/transformers/test_spectral_hybrid_lora.py +++ b/tests/transformers/test_spectral_hybrid_lora.py @@ -77,27 +77,6 @@ def _register_hybrid(manager, hybrid, adapter_name, config): hybrid.register_adapter(adapter_name) -@pytest.mark.parametrize('bind_device,expects_device_id', [(None, True), (lambda _backend: False, False)]) -def test_initialize_process_group_preserves_backend_device_binding(monkeypatch, bind_device, expects_device_id): - import torch.distributed as dist - from twinkle import Platform, torch_util - from twinkle.model.base import initialize_process_group - - calls = [] - monkeypatch.setattr(dist, 'is_initialized', lambda: False) - monkeypatch.setattr(dist, 'init_process_group', lambda **kwargs: calls.append(kwargs)) - monkeypatch.setattr(Platform, 'get_world_size', lambda: 2) - monkeypatch.setattr(Platform, 'get_rank', lambda: 0) - monkeypatch.setattr(Platform, 'get_local_device', lambda: 'cpu') - monkeypatch.setattr(Platform, 'device_backend', lambda: 'nccl') - monkeypatch.setattr(torch_util, 'set_device', lambda: None) - - initialize_process_group(bind_device) - - assert len(calls) == 1 - assert ('device_id' in calls[0]) is expects_device_id - - def test_spectral_metrics_match_weighted_formula(): singular_values = torch.tensor([4.0, 2.0, 1.0, 0.5], dtype=torch.float64) metrics = compute_spectral_metrics(singular_values, r=1)