bergson 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- bergson/__init__.py +21 -0
- bergson/__main__.py +12 -0
- bergson/attributor.py +159 -0
- bergson/build.py +261 -0
- bergson/collection.py +214 -0
- bergson/data.py +473 -0
- bergson/faiss_index.py +265 -0
- bergson/gradcheck.py +114 -0
- bergson/gradients.py +549 -0
- bergson/huggingface.py +384 -0
- bergson/math.py +102 -0
- bergson/peft.py +38 -0
- bergson/plot_eta.py +29 -0
- bergson/tmp.py +11 -0
- bergson/utils.py +60 -0
- bergson-0.0.1.dist-info/METADATA +158 -0
- bergson-0.0.1.dist-info/RECORD +21 -0
- bergson-0.0.1.dist-info/WHEEL +5 -0
- bergson-0.0.1.dist-info/entry_points.txt +2 -0
- bergson-0.0.1.dist-info/licenses/LICENSE +21 -0
- bergson-0.0.1.dist-info/top_level.txt +1 -0
bergson/__init__.py
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
__version__ = "0.0.1"
|
|
2
|
+
|
|
3
|
+
from .attributor import Attributor
|
|
4
|
+
from .collection import collect_gradients
|
|
5
|
+
from .data import DataConfig, IndexConfig, load_gradients
|
|
6
|
+
from .faiss_index import FaissConfig
|
|
7
|
+
from .gradcheck import FiniteDiff
|
|
8
|
+
from .gradients import GradientCollector, GradientProcessor, HeadConfig
|
|
9
|
+
|
|
10
|
+
__all__ = [
|
|
11
|
+
"collect_gradients",
|
|
12
|
+
"load_gradients",
|
|
13
|
+
"Attributor",
|
|
14
|
+
"FaissConfig",
|
|
15
|
+
"FiniteDiff",
|
|
16
|
+
"GradientCollector",
|
|
17
|
+
"GradientProcessor",
|
|
18
|
+
"IndexConfig",
|
|
19
|
+
"DataConfig",
|
|
20
|
+
"HeadConfig",
|
|
21
|
+
]
|
bergson/__main__.py
ADDED
bergson/attributor.py
ADDED
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
from collections import defaultdict
|
|
2
|
+
from contextlib import contextmanager
|
|
3
|
+
from typing import Generator
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
from torch import Tensor, nn
|
|
7
|
+
|
|
8
|
+
from .data import load_gradients
|
|
9
|
+
from .faiss_index import FaissConfig, FaissIndex
|
|
10
|
+
from .gradients import GradientCollector, GradientProcessor
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class TraceResult:
|
|
14
|
+
"""Result of a .trace() call."""
|
|
15
|
+
|
|
16
|
+
def __init__(self):
|
|
17
|
+
# Should be set by the Attributor after a search
|
|
18
|
+
self._indices: Tensor | None = None
|
|
19
|
+
self._scores: Tensor | None = None
|
|
20
|
+
|
|
21
|
+
@property
|
|
22
|
+
def indices(self) -> Tensor:
|
|
23
|
+
"""The indices of the top-k examples."""
|
|
24
|
+
if self._indices is None:
|
|
25
|
+
raise ValueError("No indices available. Exit the context manager first.")
|
|
26
|
+
|
|
27
|
+
return self._indices
|
|
28
|
+
|
|
29
|
+
@property
|
|
30
|
+
def scores(self) -> Tensor:
|
|
31
|
+
"""The attribution scores of the top-k examples."""
|
|
32
|
+
if self._scores is None:
|
|
33
|
+
raise ValueError("No scores available. Exit the context manager first.")
|
|
34
|
+
|
|
35
|
+
return self._scores
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class Attributor:
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
index_path: str,
|
|
42
|
+
device: str = "cpu",
|
|
43
|
+
dtype: torch.dtype = torch.float32,
|
|
44
|
+
unit_norm: bool = False,
|
|
45
|
+
faiss_cfg: FaissConfig | None = None,
|
|
46
|
+
):
|
|
47
|
+
self.device = device
|
|
48
|
+
self.dtype = dtype
|
|
49
|
+
self.unit_norm = unit_norm
|
|
50
|
+
self.faiss_index = None
|
|
51
|
+
|
|
52
|
+
# Load the gradient processor
|
|
53
|
+
self.processor = GradientProcessor.load(index_path, map_location=device)
|
|
54
|
+
|
|
55
|
+
# Load the gradient index
|
|
56
|
+
if faiss_cfg:
|
|
57
|
+
self.faiss_index = FaissIndex(index_path, faiss_cfg, device, unit_norm)
|
|
58
|
+
self.N = self.faiss_index.ntotal
|
|
59
|
+
else:
|
|
60
|
+
mmap = load_gradients(index_path)
|
|
61
|
+
|
|
62
|
+
# Copy gradients into device memory
|
|
63
|
+
self.grads = {
|
|
64
|
+
name: torch.tensor(mmap[name], device=device, dtype=dtype)
|
|
65
|
+
for name in mmap.dtype.names
|
|
66
|
+
}
|
|
67
|
+
self.N = mmap[mmap.dtype.names[0]].shape[0]
|
|
68
|
+
|
|
69
|
+
if unit_norm:
|
|
70
|
+
norm = torch.cat([grad for grad in self.grads.values()], dim=1).norm(
|
|
71
|
+
dim=1, keepdim=True
|
|
72
|
+
)
|
|
73
|
+
for name in self.grads:
|
|
74
|
+
self.grads[name] /= norm
|
|
75
|
+
|
|
76
|
+
def search(
|
|
77
|
+
self, queries: dict[str, Tensor], k: int, modules: list[str] | None = None
|
|
78
|
+
) -> tuple[Tensor, Tensor]:
|
|
79
|
+
"""
|
|
80
|
+
Search for the `k` nearest examples in the index based on the query or queries.
|
|
81
|
+
|
|
82
|
+
Args:
|
|
83
|
+
queries: The query tensor of shape [..., d].
|
|
84
|
+
k: The number of nearest examples to return for each query.
|
|
85
|
+
module: The name of the module to search for. If `None`,
|
|
86
|
+
all modules will be searched.
|
|
87
|
+
|
|
88
|
+
Returns:
|
|
89
|
+
A namedtuple containing the top `k` indices and inner products for each
|
|
90
|
+
query. Both have shape [..., k].
|
|
91
|
+
"""
|
|
92
|
+
q = {name: item.to(self.device, self.dtype) for name, item in queries.items()}
|
|
93
|
+
|
|
94
|
+
if self.unit_norm:
|
|
95
|
+
norm = torch.cat(list(q.values()), dim=1).norm(dim=1, keepdim=True)
|
|
96
|
+
|
|
97
|
+
for name in q:
|
|
98
|
+
q[name] /= norm + 1e-8
|
|
99
|
+
|
|
100
|
+
if self.faiss_index:
|
|
101
|
+
if modules:
|
|
102
|
+
raise NotImplementedError(
|
|
103
|
+
"FAISS index does not implement module-specific search."
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
q = torch.cat([q[name] for name in q], dim=1).cpu().numpy()
|
|
107
|
+
|
|
108
|
+
distances, indices = self.faiss_index.search(q, k)
|
|
109
|
+
|
|
110
|
+
return torch.from_numpy(distances.squeeze()), torch.from_numpy(
|
|
111
|
+
indices.squeeze()
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
modules = modules or list(q.keys())
|
|
115
|
+
k = min(k, self.N)
|
|
116
|
+
|
|
117
|
+
scores = torch.stack(
|
|
118
|
+
[q[name] @ self.grads[name].mT for name in modules], dim=-1
|
|
119
|
+
).sum(-1)
|
|
120
|
+
|
|
121
|
+
return torch.topk(scores, k)
|
|
122
|
+
|
|
123
|
+
@contextmanager
|
|
124
|
+
def trace(
|
|
125
|
+
self, module: nn.Module, k: int, *, precondition: bool = False
|
|
126
|
+
) -> Generator[TraceResult, None, None]:
|
|
127
|
+
"""
|
|
128
|
+
Context manager to trace the gradients of a module and return the
|
|
129
|
+
corresponding Attributor instance.
|
|
130
|
+
"""
|
|
131
|
+
mod_grads = defaultdict(list)
|
|
132
|
+
result = TraceResult()
|
|
133
|
+
|
|
134
|
+
def callback(name: str, g: Tensor):
|
|
135
|
+
# Precondition the gradient using Cholesky solve
|
|
136
|
+
if precondition:
|
|
137
|
+
eigval, eigvec = self.processor.preconditioners_eigen[name]
|
|
138
|
+
eigval_inverse_sqrt = 1.0 / (eigval).sqrt()
|
|
139
|
+
P = eigvec * eigval_inverse_sqrt @ eigvec.mT
|
|
140
|
+
g = g.flatten(1).type_as(P)
|
|
141
|
+
g = g @ P
|
|
142
|
+
else:
|
|
143
|
+
g = g.flatten(1)
|
|
144
|
+
|
|
145
|
+
# Store the gradient for later use
|
|
146
|
+
mod_grads[name].append(g.to(self.device, self.dtype, non_blocking=True))
|
|
147
|
+
|
|
148
|
+
with GradientCollector(module, callback, self.processor):
|
|
149
|
+
yield result
|
|
150
|
+
|
|
151
|
+
if not mod_grads:
|
|
152
|
+
raise ValueError("No grads collected. Did you forget to call backward?")
|
|
153
|
+
|
|
154
|
+
queries = {name: torch.cat(g, dim=1) for name, g in mod_grads.items()}
|
|
155
|
+
|
|
156
|
+
if any(q.isnan().any() for q in queries.values()):
|
|
157
|
+
raise ValueError("NaN found in queries.")
|
|
158
|
+
|
|
159
|
+
result._scores, result._indices = self.search(queries, k)
|
bergson/build.py
ADDED
|
@@ -0,0 +1,261 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import socket
|
|
3
|
+
from datetime import timedelta
|
|
4
|
+
from typing import cast
|
|
5
|
+
|
|
6
|
+
import pandas as pd
|
|
7
|
+
import torch
|
|
8
|
+
import torch.distributed as dist
|
|
9
|
+
import torch.multiprocessing as mp
|
|
10
|
+
from datasets import Dataset, IterableDataset
|
|
11
|
+
from peft import PeftConfig, PeftModel
|
|
12
|
+
from torch.distributed.elastic.multiprocessing import DefaultLogsSpecs, start_processes
|
|
13
|
+
from torch.distributed.fsdp import fully_shard
|
|
14
|
+
from tqdm.auto import tqdm
|
|
15
|
+
from transformers import (
|
|
16
|
+
AutoModelForCausalLM,
|
|
17
|
+
AutoTokenizer,
|
|
18
|
+
BitsAndBytesConfig,
|
|
19
|
+
PreTrainedModel,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
from .collection import collect_gradients
|
|
23
|
+
from .data import IndexConfig, allocate_batches, load_data_string, tokenize
|
|
24
|
+
from .gradients import GradientProcessor
|
|
25
|
+
from .peft import detect_peft_modules
|
|
26
|
+
from .utils import assert_type, get_layer_list
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def worker(rank: int, world_size: int, cfg: IndexConfig, ds: Dataset | IterableDataset):
|
|
30
|
+
torch.cuda.set_device(rank)
|
|
31
|
+
|
|
32
|
+
# These should be set by the main process
|
|
33
|
+
if world_size > 1:
|
|
34
|
+
addr = os.environ.get("MASTER_ADDR", "localhost")
|
|
35
|
+
port = os.environ.get("MASTER_PORT", "29500")
|
|
36
|
+
|
|
37
|
+
dist.init_process_group(
|
|
38
|
+
"nccl",
|
|
39
|
+
init_method=f"tcp://{addr}:{port}",
|
|
40
|
+
device_id=torch.device(f"cuda:{rank}"),
|
|
41
|
+
rank=rank,
|
|
42
|
+
timeout=timedelta(hours=1),
|
|
43
|
+
world_size=world_size,
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
match cfg.precision:
|
|
47
|
+
case "bf16":
|
|
48
|
+
dtype = torch.bfloat16
|
|
49
|
+
case "fp16":
|
|
50
|
+
dtype = torch.float16
|
|
51
|
+
case "fp32":
|
|
52
|
+
dtype = torch.float32
|
|
53
|
+
case "int4" | "int8":
|
|
54
|
+
dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
|
|
55
|
+
case "auto":
|
|
56
|
+
dtype = "auto"
|
|
57
|
+
case other:
|
|
58
|
+
raise ValueError(f"Unsupported precision: {other}")
|
|
59
|
+
|
|
60
|
+
device_map = {"": f"cuda:{rank}"} if not cfg.fsdp else "cpu"
|
|
61
|
+
quantization_config = None
|
|
62
|
+
if cfg.precision in ("int4", "int8"):
|
|
63
|
+
quantization_config = BitsAndBytesConfig(
|
|
64
|
+
load_in_4bit=cfg.precision == "int4",
|
|
65
|
+
load_in_8bit=cfg.precision == "int8",
|
|
66
|
+
bnb_4bit_compute_dtype=dtype,
|
|
67
|
+
bnb_4bit_quant_storage=dtype,
|
|
68
|
+
bnb_4bit_quant_type="nf4",
|
|
69
|
+
bnb_4bit_use_double_quant=True,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
# Try to detect PEFT model
|
|
73
|
+
try:
|
|
74
|
+
peft_config = PeftConfig.from_pretrained(cfg.model)
|
|
75
|
+
except ValueError:
|
|
76
|
+
peft_config = None
|
|
77
|
+
|
|
78
|
+
if peft_config is None:
|
|
79
|
+
# Load regular model
|
|
80
|
+
model = AutoModelForCausalLM.from_pretrained(
|
|
81
|
+
cfg.model,
|
|
82
|
+
device_map=device_map,
|
|
83
|
+
quantization_config=quantization_config,
|
|
84
|
+
dtype=dtype,
|
|
85
|
+
revision=cfg.revision,
|
|
86
|
+
)
|
|
87
|
+
target_modules = None
|
|
88
|
+
|
|
89
|
+
else:
|
|
90
|
+
# Load PEFT model
|
|
91
|
+
base_model = AutoModelForCausalLM.from_pretrained(
|
|
92
|
+
peft_config.base_model_name_or_path, # type: ignore
|
|
93
|
+
device_map=device_map,
|
|
94
|
+
quantization_config=quantization_config,
|
|
95
|
+
dtype=dtype,
|
|
96
|
+
revision=cfg.revision,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
model = PeftModel.from_pretrained(
|
|
100
|
+
base_model,
|
|
101
|
+
cfg.model,
|
|
102
|
+
device_map=device_map,
|
|
103
|
+
autocast_adapter_dtype=False,
|
|
104
|
+
)
|
|
105
|
+
target_modules = detect_peft_modules(model)
|
|
106
|
+
|
|
107
|
+
# Hack for type checking
|
|
108
|
+
model = cast(PreTrainedModel, model)
|
|
109
|
+
|
|
110
|
+
if rank == 0:
|
|
111
|
+
print(f"Model loaded with dtype: {model.dtype}")
|
|
112
|
+
|
|
113
|
+
embed = model.get_input_embeddings()
|
|
114
|
+
model.requires_grad_(False) # Freeze the model
|
|
115
|
+
embed.requires_grad_(True) # Make sure backward hooks are called though
|
|
116
|
+
|
|
117
|
+
if cfg.fsdp:
|
|
118
|
+
# Shard each individual transformer layer
|
|
119
|
+
for layer in get_layer_list(model):
|
|
120
|
+
fully_shard(layer)
|
|
121
|
+
|
|
122
|
+
# Shard the entire model
|
|
123
|
+
fully_shard(model)
|
|
124
|
+
|
|
125
|
+
if os.path.exists(cfg.processor_path):
|
|
126
|
+
if rank == 0:
|
|
127
|
+
print(f"Loading processor from '{cfg.processor_path}'")
|
|
128
|
+
|
|
129
|
+
processor = GradientProcessor.load(
|
|
130
|
+
cfg.processor_path,
|
|
131
|
+
map_location=f"cuda:{rank}",
|
|
132
|
+
)
|
|
133
|
+
else:
|
|
134
|
+
processor = GradientProcessor(
|
|
135
|
+
{},
|
|
136
|
+
projection_dim=cfg.projection_dim or None,
|
|
137
|
+
reshape_to_square=cfg.reshape_to_square,
|
|
138
|
+
projection_type=cfg.projection_type,
|
|
139
|
+
)
|
|
140
|
+
if rank == 0:
|
|
141
|
+
processor.save(cfg.run_path)
|
|
142
|
+
|
|
143
|
+
if isinstance(ds, Dataset):
|
|
144
|
+
batches = allocate_batches(ds["length"][:], cfg.token_batch_size)
|
|
145
|
+
collect_gradients(
|
|
146
|
+
model,
|
|
147
|
+
ds,
|
|
148
|
+
processor,
|
|
149
|
+
cfg.run_path,
|
|
150
|
+
batches=batches,
|
|
151
|
+
kl_divergence=cfg.loss_fn == "kl",
|
|
152
|
+
loss_reduction=cfg.loss_reduction,
|
|
153
|
+
skip_preconditioners=cfg.skip_preconditioners,
|
|
154
|
+
target_modules=target_modules,
|
|
155
|
+
head_cfgs=cfg.head_cfgs,
|
|
156
|
+
)
|
|
157
|
+
else:
|
|
158
|
+
# Convert each shard to a Dataset then collect its gradients
|
|
159
|
+
buf, shard_id = [], 0
|
|
160
|
+
|
|
161
|
+
def flush():
|
|
162
|
+
nonlocal buf, shard_id
|
|
163
|
+
if not buf:
|
|
164
|
+
return
|
|
165
|
+
ds_shard = assert_type(Dataset, Dataset.from_list(buf))
|
|
166
|
+
batches = allocate_batches(ds_shard["length"][:], cfg.token_batch_size)
|
|
167
|
+
collect_gradients(
|
|
168
|
+
model,
|
|
169
|
+
ds_shard,
|
|
170
|
+
processor,
|
|
171
|
+
os.path.join(cfg.run_path, f"shard-{shard_id:05d}"),
|
|
172
|
+
batches=batches,
|
|
173
|
+
kl_divergence=cfg.loss_fn == "kl",
|
|
174
|
+
loss_reduction=cfg.loss_reduction,
|
|
175
|
+
skip_preconditioners=cfg.skip_preconditioners,
|
|
176
|
+
target_modules=target_modules,
|
|
177
|
+
head_cfgs=cfg.head_cfgs,
|
|
178
|
+
)
|
|
179
|
+
buf.clear()
|
|
180
|
+
shard_id += 1
|
|
181
|
+
|
|
182
|
+
for ex in tqdm(ds, desc="Collecting gradients"):
|
|
183
|
+
buf.append(ex)
|
|
184
|
+
if len(buf) == cfg.stream_shard_size:
|
|
185
|
+
flush()
|
|
186
|
+
flush()
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def dist_worker(rank: int, world_size: int, cfg: IndexConfig, ds: Dataset):
|
|
190
|
+
try:
|
|
191
|
+
worker(rank, world_size, cfg, ds)
|
|
192
|
+
finally:
|
|
193
|
+
dist.destroy_process_group()
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
def estimate_advantage(ds: Dataset, cfg: IndexConfig):
|
|
197
|
+
"""Group rollouts by prompt and estimate advantages."""
|
|
198
|
+
assert isinstance(ds, Dataset), "Dataset required for advantage estimation"
|
|
199
|
+
|
|
200
|
+
df = ds.select_columns([cfg.data.prompt_column, cfg.data.reward_column]).to_pandas()
|
|
201
|
+
df = assert_type(pd.DataFrame, df)
|
|
202
|
+
|
|
203
|
+
advantages = df[cfg.data.reward_column] - df.groupby(cfg.data.prompt_column)[
|
|
204
|
+
cfg.data.reward_column
|
|
205
|
+
].transform("mean")
|
|
206
|
+
|
|
207
|
+
return advantages.tolist()
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
def build_gradient_dataset(cfg: IndexConfig):
|
|
211
|
+
# In many cases the token_batch_size may be smaller than the max length allowed by
|
|
212
|
+
# the model. If cfg.data.truncation is True, we use the tokenizer to truncate
|
|
213
|
+
tokenizer = AutoTokenizer.from_pretrained(cfg.model, revision=cfg.revision)
|
|
214
|
+
tokenizer.model_max_length = min(tokenizer.model_max_length, cfg.token_batch_size)
|
|
215
|
+
|
|
216
|
+
# Do all the data loading and preprocessing on the main process
|
|
217
|
+
ds = load_data_string(cfg.data.dataset, cfg.data.split, streaming=cfg.streaming)
|
|
218
|
+
|
|
219
|
+
remove_columns = ds.column_names if cfg.drop_columns else None
|
|
220
|
+
ds = ds.map(
|
|
221
|
+
tokenize,
|
|
222
|
+
batched=True,
|
|
223
|
+
fn_kwargs=dict(args=cfg.data, tokenizer=tokenizer),
|
|
224
|
+
remove_columns=remove_columns,
|
|
225
|
+
)
|
|
226
|
+
if cfg.data.reward_column:
|
|
227
|
+
ds = ds.add_column(
|
|
228
|
+
"advantage",
|
|
229
|
+
estimate_advantage(ds, cfg),
|
|
230
|
+
new_fingerprint="advantage", # type: ignore
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
world_size = torch.cuda.device_count()
|
|
234
|
+
if world_size <= 1:
|
|
235
|
+
# Run the worker directly if no distributed training is needed. This is great
|
|
236
|
+
# for debugging purposes.
|
|
237
|
+
worker(0, 1, cfg, ds)
|
|
238
|
+
else:
|
|
239
|
+
# Set up multiprocessing and distributed training
|
|
240
|
+
mp.set_sharing_strategy("file_system")
|
|
241
|
+
|
|
242
|
+
# Find an available port for distributed training
|
|
243
|
+
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
|
244
|
+
s.bind(("", 0))
|
|
245
|
+
_, port = s.getsockname()
|
|
246
|
+
|
|
247
|
+
ctx = start_processes(
|
|
248
|
+
"build",
|
|
249
|
+
dist_worker,
|
|
250
|
+
args={i: (i, world_size, cfg, ds) for i in range(world_size)},
|
|
251
|
+
envs={
|
|
252
|
+
i: {
|
|
253
|
+
"LOCAL_RANK": str(i),
|
|
254
|
+
"MASTER_ADDR": "localhost",
|
|
255
|
+
"MASTER_PORT": str(port),
|
|
256
|
+
}
|
|
257
|
+
for i in range(world_size)
|
|
258
|
+
},
|
|
259
|
+
logs_specs=DefaultLogsSpecs(),
|
|
260
|
+
)
|
|
261
|
+
ctx.wait()
|
bergson/collection.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
import math
|
|
2
|
+
from typing import Literal
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import torch
|
|
6
|
+
import torch.distributed as dist
|
|
7
|
+
import torch.nn.functional as F
|
|
8
|
+
from datasets import Dataset, Value
|
|
9
|
+
from tqdm.auto import tqdm
|
|
10
|
+
from transformers import PreTrainedModel
|
|
11
|
+
|
|
12
|
+
from .data import create_index, pad_and_tensor
|
|
13
|
+
from .gradients import GradientCollector, GradientProcessor, HeadConfig
|
|
14
|
+
from .peft import set_peft_enabled
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def collect_gradients(
|
|
18
|
+
model: PreTrainedModel,
|
|
19
|
+
data: Dataset,
|
|
20
|
+
processor: GradientProcessor,
|
|
21
|
+
path: str,
|
|
22
|
+
*,
|
|
23
|
+
batches: list[list[int]] | None = None,
|
|
24
|
+
kl_divergence: bool | None = None,
|
|
25
|
+
loss_reduction: Literal["mean", "sum"] = "mean",
|
|
26
|
+
skip_preconditioners: bool = False,
|
|
27
|
+
target_modules: set[str] | None = None,
|
|
28
|
+
head_cfgs: dict[str, HeadConfig] = {},
|
|
29
|
+
):
|
|
30
|
+
"""
|
|
31
|
+
Compute projected gradients using a subset of the dataset.
|
|
32
|
+
"""
|
|
33
|
+
rank = dist.get_rank() if dist.is_initialized() else 0
|
|
34
|
+
|
|
35
|
+
# Batch size of one by default
|
|
36
|
+
if batches is None:
|
|
37
|
+
batches = [[idx] for idx in range(len(data))]
|
|
38
|
+
|
|
39
|
+
# Mutable state for the GradientCollector callback
|
|
40
|
+
mod_grads = {}
|
|
41
|
+
preconditioners = {}
|
|
42
|
+
|
|
43
|
+
# TODO: Handle this more elegantly
|
|
44
|
+
dtype = torch.float32 if model.dtype == torch.float32 else torch.float16
|
|
45
|
+
np_dtype = np.float32 if dtype == torch.float32 else np.float16
|
|
46
|
+
lo = torch.finfo(dtype).min
|
|
47
|
+
hi = torch.finfo(dtype).max
|
|
48
|
+
|
|
49
|
+
def callback(name: str, g: torch.Tensor):
|
|
50
|
+
g = g.flatten(1).clamp_(lo, hi)
|
|
51
|
+
|
|
52
|
+
# Asynchronously move the gradient to CPU and convert to fp16
|
|
53
|
+
mod_grads[name] = g.to(device="cpu", dtype=dtype, non_blocking=True)
|
|
54
|
+
|
|
55
|
+
# Compute the outer product of the flattened gradient
|
|
56
|
+
if not skip_preconditioners:
|
|
57
|
+
g = g.float()
|
|
58
|
+
preconditioner = preconditioners.get(name, None)
|
|
59
|
+
if preconditioner is None:
|
|
60
|
+
preconditioners[name] = g.mT @ g
|
|
61
|
+
else:
|
|
62
|
+
preconditioner.addmm_(g.mT, g)
|
|
63
|
+
|
|
64
|
+
collector = GradientCollector(
|
|
65
|
+
model.base_model,
|
|
66
|
+
callback,
|
|
67
|
+
processor,
|
|
68
|
+
target_modules=target_modules,
|
|
69
|
+
head_cfgs=head_cfgs,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
# Allocate space ahead of time for the gradients
|
|
73
|
+
grad_sizes = {name: math.prod(s) for name, s in collector.shapes().items()}
|
|
74
|
+
|
|
75
|
+
# Allocate structured space ahead of time for the gradients
|
|
76
|
+
grad_buffer = create_index(
|
|
77
|
+
path, num_grads=len(data), grad_sizes=grad_sizes, dtype=np_dtype
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
per_doc_losses = torch.full(
|
|
81
|
+
(len(data),),
|
|
82
|
+
device=model.device,
|
|
83
|
+
dtype=dtype,
|
|
84
|
+
fill_value=0.0,
|
|
85
|
+
)
|
|
86
|
+
|
|
87
|
+
for indices in tqdm(batches, disable=rank != 0, desc="Building index"):
|
|
88
|
+
batch = data[indices]
|
|
89
|
+
x, y = pad_and_tensor(
|
|
90
|
+
batch["input_ids"], # type: ignore
|
|
91
|
+
labels=batch.get("labels"), # type: ignore
|
|
92
|
+
device=model.device,
|
|
93
|
+
)
|
|
94
|
+
masks = y[:, 1:] != -100
|
|
95
|
+
denoms = masks.sum(dim=1, dtype=dtype) if loss_reduction == "mean" else 1.0
|
|
96
|
+
|
|
97
|
+
if kl_divergence:
|
|
98
|
+
with torch.inference_mode():
|
|
99
|
+
set_peft_enabled(model, False)
|
|
100
|
+
ref_lps = torch.log_softmax(model(x).logits[:, :-1], dim=-1)
|
|
101
|
+
set_peft_enabled(model, True)
|
|
102
|
+
|
|
103
|
+
with collector:
|
|
104
|
+
ft_lps = torch.log_softmax(model(x).logits[:, :-1], dim=-1)
|
|
105
|
+
|
|
106
|
+
# Compute average KL across all unmasked tokens
|
|
107
|
+
kls = torch.sum(ft_lps.exp() * (ft_lps - ref_lps), dim=-1)
|
|
108
|
+
losses = torch.sum(kls * masks, dim=-1) / denoms
|
|
109
|
+
if "advantage" in batch:
|
|
110
|
+
losses *= torch.tensor(batch["advantage"], device=losses.device)
|
|
111
|
+
|
|
112
|
+
losses.mean().backward()
|
|
113
|
+
else:
|
|
114
|
+
with collector:
|
|
115
|
+
logits = model(x).logits[:, :-1]
|
|
116
|
+
|
|
117
|
+
losses = F.cross_entropy(
|
|
118
|
+
logits.reshape(-1, logits.size(-1)),
|
|
119
|
+
y[:, 1:].flatten(),
|
|
120
|
+
reduction="none",
|
|
121
|
+
).reshape_as(y[:, 1:])
|
|
122
|
+
losses = losses.sum(1) / denoms
|
|
123
|
+
if "advantage" in batch:
|
|
124
|
+
losses *= torch.tensor(batch["advantage"], device=losses.device)
|
|
125
|
+
|
|
126
|
+
losses.mean().backward()
|
|
127
|
+
|
|
128
|
+
# Weirdly you need to explicitly synchronize here in order to make sure that
|
|
129
|
+
# the nonblocking copies actually finish before we call .numpy()
|
|
130
|
+
model.zero_grad()
|
|
131
|
+
torch.cuda.synchronize()
|
|
132
|
+
|
|
133
|
+
# It turns out that it's very important for efficiency to write the gradients
|
|
134
|
+
# sequentially instead of first concatenating them, then writing to one vector
|
|
135
|
+
for module_name in mod_grads.keys():
|
|
136
|
+
grad_buffer[module_name][indices] = mod_grads[module_name].numpy()
|
|
137
|
+
|
|
138
|
+
mod_grads.clear()
|
|
139
|
+
per_doc_losses[indices] = losses.detach().type_as(per_doc_losses)
|
|
140
|
+
|
|
141
|
+
process_preconditioners(processor, preconditioners, len(data))
|
|
142
|
+
|
|
143
|
+
if dist.is_initialized():
|
|
144
|
+
dist.reduce(per_doc_losses, dst=0)
|
|
145
|
+
|
|
146
|
+
if rank == 0:
|
|
147
|
+
data = data.add_column(
|
|
148
|
+
"loss",
|
|
149
|
+
per_doc_losses.cpu().numpy(),
|
|
150
|
+
feature=Value("float16" if dtype == torch.float16 else "float32"),
|
|
151
|
+
new_fingerprint="loss",
|
|
152
|
+
)
|
|
153
|
+
data.save_to_disk(path + "/data.hf")
|
|
154
|
+
|
|
155
|
+
processor.save(path)
|
|
156
|
+
|
|
157
|
+
# Make sure the gradients are written to disk
|
|
158
|
+
grad_buffer.flush()
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def process_preconditioners(
|
|
162
|
+
processor: GradientProcessor,
|
|
163
|
+
preconditioners: dict[str, torch.Tensor],
|
|
164
|
+
len_data: int,
|
|
165
|
+
):
|
|
166
|
+
"""
|
|
167
|
+
Aggregate preconditioners across ranks and compute their eigen decomposition
|
|
168
|
+
distributed across all ranks.
|
|
169
|
+
"""
|
|
170
|
+
|
|
171
|
+
rank = dist.get_rank() if dist.is_initialized() else 0
|
|
172
|
+
world_size = dist.get_world_size() if dist.is_initialized() else 1
|
|
173
|
+
preconditioners_eigen = {}
|
|
174
|
+
if rank == 0:
|
|
175
|
+
print("Saving preconditioners...")
|
|
176
|
+
for name, prec in preconditioners.items():
|
|
177
|
+
if dist.is_initialized():
|
|
178
|
+
dist.all_reduce(prec)
|
|
179
|
+
|
|
180
|
+
preconditioners[name] = prec / len_data
|
|
181
|
+
|
|
182
|
+
processor.preconditioners = preconditioners
|
|
183
|
+
|
|
184
|
+
if rank == 0:
|
|
185
|
+
print("Computing preconditioner eigen decompositions...")
|
|
186
|
+
names = list(preconditioners.keys())
|
|
187
|
+
names_per_rank = names[rank::world_size]
|
|
188
|
+
|
|
189
|
+
for name in names_per_rank:
|
|
190
|
+
original_dtype = preconditioners[name].dtype
|
|
191
|
+
prec = preconditioners[name].to(dtype=torch.float64)
|
|
192
|
+
eigvals, eigvecs = torch.linalg.eigh(prec)
|
|
193
|
+
preconditioners_eigen[name] = (
|
|
194
|
+
eigvals.to(dtype=original_dtype).contiguous(),
|
|
195
|
+
eigvecs.to(dtype=original_dtype).contiguous(),
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
if rank == 0:
|
|
199
|
+
print("Gathering and saving preconditioner eigen decompositions...")
|
|
200
|
+
|
|
201
|
+
for name in names:
|
|
202
|
+
prec = preconditioners[name]
|
|
203
|
+
if name not in preconditioners_eigen:
|
|
204
|
+
eigval = torch.zeros(prec.size(0), dtype=prec.dtype, device=prec.device)
|
|
205
|
+
eigvec = torch.zeros_like(prec)
|
|
206
|
+
else:
|
|
207
|
+
eigval, eigvec = preconditioners_eigen[name]
|
|
208
|
+
|
|
209
|
+
dist.all_reduce(eigval, op=dist.ReduceOp.SUM) if dist.is_initialized() else None
|
|
210
|
+
dist.all_reduce(eigvec, op=dist.ReduceOp.SUM) if dist.is_initialized() else None
|
|
211
|
+
|
|
212
|
+
preconditioners_eigen[name] = (eigval, eigvec)
|
|
213
|
+
if rank == 0:
|
|
214
|
+
processor.preconditioners_eigen = preconditioners_eigen
|