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 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
@@ -0,0 +1,12 @@
1
+ from simple_parsing import parse
2
+
3
+ from .build import build_gradient_dataset
4
+ from .data import IndexConfig
5
+
6
+
7
+ def main():
8
+ build_gradient_dataset(parse(IndexConfig))
9
+
10
+
11
+ if __name__ == "__main__":
12
+ main()
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