Skip to content

Commit f433265

Browse files
committed
fix device
1 parent b88717d commit f433265

2 files changed

Lines changed: 24 additions & 20 deletions

File tree

bergson/collector/gradient_collectors.py

Lines changed: 22 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -84,26 +84,28 @@ def setup(self) -> None:
8484
8585
Sets up a Builder for gradient storage if not using a Scorer.
8686
"""
87+
model_device = (
88+
getattr(self.model, "device", None) or next(self.model.parameters()).device
89+
)
90+
model_dtype = (
91+
getattr(self.model, "dtype", None) or next(self.model.parameters()).dtype
92+
)
93+
8794
assert isinstance(
88-
self.model.device, torch.device
95+
model_device, torch.device
8996
), "Model device is not set correctly"
90-
if self.cfg.include_bias and self.processor.normalizers is not None:
91-
raise NotImplementedError(
92-
"Bias with normalizers not supported yet, "
93-
"consider disabling bias inclusion for now."
94-
)
9597

9698
# TODO: handle more elegantly?
9799
self.save_dtype = (
98-
torch.float32 if self.model.dtype == torch.float32 else torch.float16
100+
torch.float32 if model_dtype == torch.float32 else torch.float16
99101
)
100102

101103
self.lo = torch.finfo(self.save_dtype).min
102104
self.hi = torch.finfo(self.save_dtype).max
103105

104106
self.per_doc_losses = torch.full(
105107
(len(self.data),),
106-
device=self.model.device,
108+
device=model_device,
107109
dtype=self.save_dtype,
108110
fill_value=0.0,
109111
)
@@ -363,9 +365,14 @@ class TraceCollector(HookCollectorBase):
363365
"""Dtype for stored gradients."""
364366

365367
def setup(self) -> None:
368+
369+
model_dtype = (
370+
getattr(self.model, "dtype", None) or next(self.model.parameters()).dtype
371+
)
372+
366373
# TODO: handle more elegantly?
367374
self.save_dtype = (
368-
torch.float32 if self.model.dtype == torch.float32 else torch.float16
375+
torch.float32 if model_dtype == torch.float32 else torch.float16
369376
)
370377

371378
self.lo = torch.finfo(self.save_dtype).min
@@ -473,9 +480,14 @@ class StreamingGradientCollector(HookCollectorBase):
473480
"""Dtype for stored gradients."""
474481

475482
def setup(self) -> None:
483+
484+
model_dtype = (
485+
getattr(self.model, "dtype", None) or next(self.model.parameters()).dtype
486+
)
487+
476488
# TODO: handle more elegantly?
477489
self.save_dtype = (
478-
torch.float32 if self.model.dtype == torch.float32 else torch.float16
490+
torch.float32 if model_dtype == torch.float32 else torch.float16
479491
)
480492

481493
self.lo = torch.finfo(self.save_dtype).min

tests/test_gradients.py

Lines changed: 2 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -37,8 +37,8 @@ def test_params():
3737
def simple_model_class(test_params):
3838
"""Factory for creating test model classes.
3939
40-
Creates simple neural network models with device/dtype properties required
41-
by GradientCollector. Supports both single-layer and two-layer architectures.
40+
Creates simple neural network models for testing gradient collection.
41+
Supports both single-layer and two-layer architectures.
4242
4343
Returns:
4444
callable: Factory function that takes:
@@ -69,14 +69,6 @@ def __init__(self):
6969
def forward(self, x):
7070
return self.layers(x)
7171

72-
@property
73-
def device(self):
74-
return next(self.parameters()).device
75-
76-
@property
77-
def dtype(self):
78-
return next(self.parameters()).dtype
79-
8072
return SimpleModel
8173

8274
return _make_model

0 commit comments

Comments
 (0)