1- """Multi-LoRA adapter and config dataclasses.
2-
3- See docs/multi-lora-implementation-plan.md §1.2. These are plain, frozen
4- dataclasses — no NeMo-RL or torch imports here so they stay cheap to load.
5- """
1+ """Configuration consumed by the packed multi-adapter SFT path."""
62
73from __future__ import annotations
84
95from dataclasses import dataclass , field
10- from typing import Literal
116
127
138@dataclass (frozen = True )
149class MultiLoRAAdapter :
1510 """A single LoRA adapter slot in a multi-adapter run."""
1611
1712 name : str
18- lora_cfg : dict # NeMo-RL lora_cfg block, verbatim
19- optimizer : dict # NeMo-RL optimizer block, verbatim
20- data : dict # NeMo-RL data block, verbatim
21- weight : float = 1.0
22- nemo_gym : str | None = None # reserved for Phase 3 (RL)
23- pin : Literal ["gpu" , "cpu" , "auto" ] = "auto"
13+ data : dict
2414
2515
2616@dataclass (frozen = True )
2717class MultiLoRAConfig :
28- """Top-level multi-LoRA config block ."""
18+ """The fields that the native concurrent SFT implementation reads ."""
2919
3020 enabled : bool
31- schedule : Literal ["round_robin" , "weighted" , "progress_aware" ] = "round_robin"
32- global_batch_size : int = 32
33- # Number of consecutive steps the trainer should stay on a single adapter
34- # (or group, in concurrent mode) before re-querying the scheduler. Reduces
35- # weight-swap pressure (sequential mode) or group churn (concurrent mode).
36- # K=1 (default) = re-query every step.
37- steps_per_adapter : int = 1
38- # Where inactive-adapter state lives. "cpu" offloads (saves GPU memory,
39- # adds a copy per swap); "cuda" keeps every adapter resident on GPU
40- # (zero swap latency, ~tens of MB per adapter at rank=16). Default
41- # "cpu" stays safe for big models / many adapters.
42- storage_device : Literal ["cpu" , "cuda" ] = "cpu"
43- # Execution mode. "sequential" (default) = Phase A weight-swap loop. "concurrent"
44- # = Phase B token-packed forward (one base pass serves all GPU-pinned adapters).
45- execution_mode : Literal ["sequential" , "concurrent" ] = "sequential"
46- # HBM budget for the memory planner (concurrent mode only). ``None`` = auto-detect
47- # from torch.cuda.mem_get_info at runtime.
48- hbm_budget_gb : float | None = None
49- # Per-adapter micro batch size for the NeMo-RL multi-adapter SFT path
50- # (round-robin packer in :mod:`nemo_rl.models.multi_lora.data`). The
51- # global packed batch is ``len(adapters) * batch_size_per_adapter`` rows.
5221 batch_size_per_adapter : int = 1
5322 adapters : list [MultiLoRAAdapter ] = field (default_factory = list )
5423
@@ -57,12 +26,6 @@ def from_dict(cls, d: dict) -> "MultiLoRAConfig":
5726 adapters = [MultiLoRAAdapter (** a ) for a in d .get ("adapters" , [])]
5827 return cls (
5928 enabled = d .get ("enabled" , False ),
60- schedule = d .get ("schedule" , "round_robin" ),
61- global_batch_size = d .get ("global_batch_size" , 32 ),
62- steps_per_adapter = d .get ("steps_per_adapter" , 1 ),
63- storage_device = d .get ("storage_device" , "cpu" ),
64- execution_mode = d .get ("execution_mode" , "sequential" ),
65- hbm_budget_gb = d .get ("hbm_budget_gb" , None ),
6629 batch_size_per_adapter = d .get ("batch_size_per_adapter" , 1 ),
6730 adapters = adapters ,
6831 )
@@ -73,34 +36,8 @@ def validate(self) -> None:
7336 names = [a .name for a in self .adapters ]
7437 if len (names ) != len (set (names )):
7538 raise ValueError (f"Duplicate adapter names: { names } " )
76- if self .storage_device not in ("cpu" , "cuda" ):
77- raise ValueError (
78- f"multi_lora.storage_device must be 'cpu' or 'cuda', got "
79- f"{ self .storage_device !r} "
80- )
81- if self .steps_per_adapter < 1 :
82- raise ValueError (
83- f"multi_lora.steps_per_adapter must be >= 1, got "
84- f"{ self .steps_per_adapter } "
85- )
86- if self .execution_mode not in ("sequential" , "concurrent" ):
87- raise ValueError (
88- f"multi_lora.execution_mode must be 'sequential' | 'concurrent', "
89- f"got { self .execution_mode !r} "
90- )
91- if self .hbm_budget_gb is not None and self .hbm_budget_gb <= 0 :
92- raise ValueError (
93- f"multi_lora.hbm_budget_gb must be positive when set, got "
94- f"{ self .hbm_budget_gb } "
95- )
9639 if self .batch_size_per_adapter < 1 :
9740 raise ValueError (
9841 f"multi_lora.batch_size_per_adapter must be >= 1, got "
9942 f"{ self .batch_size_per_adapter } "
10043 )
101- for a in self .adapters :
102- if a .pin not in ("gpu" , "cpu" , "auto" ):
103- raise ValueError (
104- f"adapter { a .name !r} pin must be 'gpu' | 'cpu' | 'auto', "
105- f"got { a .pin !r} "
106- )
0 commit comments