graphssl 0.1.0__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.
- graphssl/__init__.py +63 -0
- graphssl/augmentation/__init__.py +11 -0
- graphssl/augmentation/compose.py +60 -0
- graphssl/augmentation/functional.py +127 -0
- graphssl/augmentation/transforms.py +84 -0
- graphssl/config/__init__.py +14 -0
- graphssl/config/load.py +114 -0
- graphssl/config/schema.py +413 -0
- graphssl/core/__init__.py +14 -0
- graphssl/core/augmentation.py +16 -0
- graphssl/core/callback.py +44 -0
- graphssl/core/encoder.py +18 -0
- graphssl/core/model.py +58 -0
- graphssl/core/registry.py +47 -0
- graphssl/data/__init__.py +1 -0
- graphssl/data/datamodule.py +100 -0
- graphssl/encoders/__init__.py +3 -0
- graphssl/encoders/gcn.py +104 -0
- graphssl/encoders/gin.py +181 -0
- graphssl/encoders/transformer.py +206 -0
- graphssl/evaluation/__init__.py +3 -0
- graphssl/evaluation/knn.py +52 -0
- graphssl/evaluation/linear_probe.py +116 -0
- graphssl/evaluation/visualization.py +206 -0
- graphssl/losses/__init__.py +6 -0
- graphssl/losses/barlow.py +35 -0
- graphssl/losses/combined.py +67 -0
- graphssl/losses/dino.py +38 -0
- graphssl/losses/nt_xent.py +40 -0
- graphssl/losses/regression.py +44 -0
- graphssl/losses/vicreg.py +56 -0
- graphssl/models/__init__.py +14 -0
- graphssl/models/afgrl.py +133 -0
- graphssl/models/barlow_twins.py +68 -0
- graphssl/models/bgrl.py +110 -0
- graphssl/models/dgi.py +85 -0
- graphssl/models/graphcl.py +68 -0
- graphssl/models/graphdino.py +182 -0
- graphssl/models/supervised.py +48 -0
- graphssl/models/vicreg.py +72 -0
- graphssl/nn/__init__.py +4 -0
- graphssl/nn/dino_head.py +99 -0
- graphssl/nn/mlp.py +52 -0
- graphssl/nn/norm.py +24 -0
- graphssl/nn/pooling.py +18 -0
- graphssl/registry/__init__.py +1 -0
- graphssl/registry/registry.py +9 -0
- graphssl/training/__init__.py +2 -0
- graphssl/training/callbacks.py +138 -0
- graphssl/training/trainer.py +108 -0
- graphssl/utils/__init__.py +3 -0
- graphssl/utils/ema.py +14 -0
- graphssl/utils/positive_miner.py +119 -0
- graphssl/utils/schedulers.py +55 -0
- graphssl-0.1.0.dist-info/METADATA +273 -0
- graphssl-0.1.0.dist-info/RECORD +59 -0
- graphssl-0.1.0.dist-info/WHEEL +5 -0
- graphssl-0.1.0.dist-info/licenses/LICENSE +21 -0
- graphssl-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,413 @@
|
|
|
1
|
+
"""Validated configuration dataclasses for all model components.
|
|
2
|
+
|
|
3
|
+
Each dataclass mirrors the corresponding YAML block and raises ValueError at
|
|
4
|
+
construction time if a value is out of range, so misconfigured runs fail at
|
|
5
|
+
load time rather than mid-training.
|
|
6
|
+
|
|
7
|
+
Usage (config-driven construction, uniform across all models)::
|
|
8
|
+
|
|
9
|
+
from graphssl.models import BGRL
|
|
10
|
+
model = BGRL(config=yaml_config["model"], in_channels=dataset.num_features)
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
from dataclasses import dataclass, field
|
|
16
|
+
from typing import Any, Dict, List, Optional
|
|
17
|
+
|
|
18
|
+
import torch.nn as nn
|
|
19
|
+
|
|
20
|
+
# ---------------------------------------------------------------------------
|
|
21
|
+
# Shared building blocks
|
|
22
|
+
# ---------------------------------------------------------------------------
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass
|
|
26
|
+
class EncoderConfig:
|
|
27
|
+
"""Validated configuration for graph encoders.
|
|
28
|
+
|
|
29
|
+
All three encoder backbones (GCN, GIN, Transformer) share this config.
|
|
30
|
+
Fields that don't apply to a given encoder are silently ignored during
|
|
31
|
+
build() via introspection — no extra per-encoder config class needed.
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
name: str
|
|
35
|
+
hidden_dim: int
|
|
36
|
+
num_layers: int
|
|
37
|
+
mlp_ratio: float = 2.0
|
|
38
|
+
drop: float = 0.2
|
|
39
|
+
pool: bool = True
|
|
40
|
+
norm_type: str = "batch" # 'batch', 'layer', or 'none'
|
|
41
|
+
edge_dim: Optional[int] = None # enables edge-feature-aware conv (GINEConv / TransformerConv)
|
|
42
|
+
node_emb_num_classes: Optional[int] = (
|
|
43
|
+
None # categorical node features (e.g. ZINC: 28 atom types)
|
|
44
|
+
)
|
|
45
|
+
edge_emb_num_classes: Optional[int] = (
|
|
46
|
+
None # categorical edge features (e.g. ZINC: 4 bond types)
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
def __post_init__(self):
|
|
50
|
+
if self.hidden_dim <= 0:
|
|
51
|
+
raise ValueError(f"hidden_dim must be > 0, got {self.hidden_dim}")
|
|
52
|
+
if self.num_layers <= 0:
|
|
53
|
+
raise ValueError(f"num_layers must be > 0, got {self.num_layers}")
|
|
54
|
+
if not (0.0 <= self.drop < 1.0):
|
|
55
|
+
raise ValueError(f"drop must be in [0, 1), got {self.drop}")
|
|
56
|
+
if self.norm_type not in ("batch", "layer", "none"):
|
|
57
|
+
raise ValueError(
|
|
58
|
+
f"norm_type must be 'batch', 'layer', or 'none', got {self.norm_type!r}"
|
|
59
|
+
)
|
|
60
|
+
if self.edge_dim is not None and self.edge_dim <= 0:
|
|
61
|
+
raise ValueError(f"edge_dim must be > 0, got {self.edge_dim}")
|
|
62
|
+
if self.node_emb_num_classes is not None and self.node_emb_num_classes <= 0:
|
|
63
|
+
raise ValueError(f"node_emb_num_classes must be > 0, got {self.node_emb_num_classes}")
|
|
64
|
+
if self.edge_emb_num_classes is not None and self.edge_emb_num_classes <= 0:
|
|
65
|
+
raise ValueError(f"edge_emb_num_classes must be > 0, got {self.edge_emb_num_classes}")
|
|
66
|
+
if self.edge_emb_num_classes is not None and self.edge_dim is None:
|
|
67
|
+
raise ValueError(
|
|
68
|
+
"edge_emb_num_classes requires edge_dim to be set "
|
|
69
|
+
"(the embedding output dimension must be defined)."
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
def build(self, in_channels: int) -> nn.Module:
|
|
73
|
+
"""Instantiate the encoder via the ENCODERS registry.
|
|
74
|
+
|
|
75
|
+
Only the kwargs accepted by the encoder's __init__ are forwarded,
|
|
76
|
+
so EncoderConfig can carry universal fields (num_layers, drop, …)
|
|
77
|
+
without breaking encoders that don't expose those parameters (e.g. GCN).
|
|
78
|
+
"""
|
|
79
|
+
import inspect
|
|
80
|
+
|
|
81
|
+
from graphssl.registry import ENCODERS
|
|
82
|
+
|
|
83
|
+
cls = ENCODERS.get_builder(self.name)
|
|
84
|
+
valid = set(inspect.signature(cls.__init__).parameters) - {"self"}
|
|
85
|
+
kwargs = {k: v for k, v in vars(self).items() if k != "name" and k in valid}
|
|
86
|
+
return ENCODERS.build(self.name, in_channels=in_channels, **kwargs)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
@dataclass
|
|
90
|
+
class HeadConfig:
|
|
91
|
+
"""Validated configuration for the DINOHead projection head."""
|
|
92
|
+
|
|
93
|
+
name: str
|
|
94
|
+
proj_hidden: int
|
|
95
|
+
bottleneck_dim: int
|
|
96
|
+
n_prototypes: int
|
|
97
|
+
student_temp: float = 0.1
|
|
98
|
+
teacher_temp: float = 0.04
|
|
99
|
+
center_momentum: float = 0.9
|
|
100
|
+
warmup_teacher_temp: float = 0.04
|
|
101
|
+
warmup_teacher_temp_epochs: int = 0
|
|
102
|
+
|
|
103
|
+
def __post_init__(self):
|
|
104
|
+
for fname, val in [
|
|
105
|
+
("proj_hidden", self.proj_hidden),
|
|
106
|
+
("bottleneck_dim", self.bottleneck_dim),
|
|
107
|
+
("n_prototypes", self.n_prototypes),
|
|
108
|
+
]:
|
|
109
|
+
if val <= 0:
|
|
110
|
+
raise ValueError(f"{fname} must be > 0, got {val}")
|
|
111
|
+
if self.student_temp <= 0 or self.teacher_temp <= 0:
|
|
112
|
+
raise ValueError("student_temp and teacher_temp must be > 0")
|
|
113
|
+
if self.warmup_teacher_temp <= 0:
|
|
114
|
+
raise ValueError("warmup_teacher_temp must be > 0")
|
|
115
|
+
if self.warmup_teacher_temp > self.teacher_temp:
|
|
116
|
+
raise ValueError(
|
|
117
|
+
f"warmup_teacher_temp ({self.warmup_teacher_temp}) must be "
|
|
118
|
+
f"<= teacher_temp ({self.teacher_temp})"
|
|
119
|
+
)
|
|
120
|
+
if self.warmup_teacher_temp_epochs < 0:
|
|
121
|
+
raise ValueError(
|
|
122
|
+
f"warmup_teacher_temp_epochs must be >= 0, got {self.warmup_teacher_temp_epochs}"
|
|
123
|
+
)
|
|
124
|
+
if not (0.0 <= self.center_momentum < 1.0):
|
|
125
|
+
raise ValueError(f"center_momentum must be in [0, 1), got {self.center_momentum}")
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
@dataclass
|
|
129
|
+
class AugmentConfig:
|
|
130
|
+
"""Name and keyword arguments for a single augmentation step."""
|
|
131
|
+
|
|
132
|
+
name: str
|
|
133
|
+
kwargs: Dict[str, Any] = field(default_factory=dict)
|
|
134
|
+
|
|
135
|
+
def __post_init__(self):
|
|
136
|
+
if not self.name:
|
|
137
|
+
raise ValueError("Augmentation name cannot be empty")
|
|
138
|
+
|
|
139
|
+
@classmethod
|
|
140
|
+
def from_dict(cls, d: dict) -> AugmentConfig:
|
|
141
|
+
d = dict(d)
|
|
142
|
+
name = d.pop("name")
|
|
143
|
+
return cls(name=name, kwargs=d)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
# ---------------------------------------------------------------------------
|
|
147
|
+
# Per-model configs
|
|
148
|
+
# ---------------------------------------------------------------------------
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
@dataclass
|
|
152
|
+
class DGIConfig:
|
|
153
|
+
"""Validated configuration for DGI."""
|
|
154
|
+
|
|
155
|
+
encoder: EncoderConfig
|
|
156
|
+
corruption: str = "shuffle_nodes"
|
|
157
|
+
shuffle_ratio: float = 1.0
|
|
158
|
+
|
|
159
|
+
def __post_init__(self):
|
|
160
|
+
if self.corruption not in ("shuffle_nodes", "shuffle_edges"):
|
|
161
|
+
raise ValueError(
|
|
162
|
+
f"corruption must be 'shuffle_nodes' or 'shuffle_edges', got {self.corruption!r}"
|
|
163
|
+
)
|
|
164
|
+
if not (0.0 < self.shuffle_ratio <= 1.0):
|
|
165
|
+
raise ValueError(f"shuffle_ratio must be in (0, 1], got {self.shuffle_ratio}")
|
|
166
|
+
|
|
167
|
+
@classmethod
|
|
168
|
+
def from_dict(cls, d: dict) -> DGIConfig:
|
|
169
|
+
return cls(
|
|
170
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
171
|
+
corruption=d.get("corruption", "shuffle_nodes"),
|
|
172
|
+
shuffle_ratio=d.get("shuffle_ratio", 1.0),
|
|
173
|
+
)
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
@dataclass
|
|
177
|
+
class GraphCLConfig:
|
|
178
|
+
"""Validated configuration for GraphCL."""
|
|
179
|
+
|
|
180
|
+
encoder: EncoderConfig
|
|
181
|
+
augment: List[AugmentConfig] = field(default_factory=list)
|
|
182
|
+
proj_dim: int = 128
|
|
183
|
+
tau: float = 0.5
|
|
184
|
+
|
|
185
|
+
def __post_init__(self):
|
|
186
|
+
if self.proj_dim <= 0:
|
|
187
|
+
raise ValueError(f"proj_dim must be > 0, got {self.proj_dim}")
|
|
188
|
+
if self.tau <= 0:
|
|
189
|
+
raise ValueError(f"tau must be > 0, got {self.tau}")
|
|
190
|
+
|
|
191
|
+
@classmethod
|
|
192
|
+
def from_dict(cls, d: dict) -> GraphCLConfig:
|
|
193
|
+
return cls(
|
|
194
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
195
|
+
augment=[AugmentConfig.from_dict(a) for a in d.get("augment", [])],
|
|
196
|
+
proj_dim=d.get("proj_dim", 128),
|
|
197
|
+
tau=d.get("tau", 0.5),
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
@dataclass
|
|
202
|
+
class VICRegConfig:
|
|
203
|
+
"""Validated configuration for VICReg."""
|
|
204
|
+
|
|
205
|
+
encoder: EncoderConfig
|
|
206
|
+
augment: List[AugmentConfig] = field(default_factory=list)
|
|
207
|
+
proj_dim: int = 256
|
|
208
|
+
invariance: float = 25.0
|
|
209
|
+
variance: float = 25.0
|
|
210
|
+
covariance: float = 1.0
|
|
211
|
+
|
|
212
|
+
def __post_init__(self):
|
|
213
|
+
if self.proj_dim <= 0:
|
|
214
|
+
raise ValueError(f"proj_dim must be > 0, got {self.proj_dim}")
|
|
215
|
+
for name, val in [
|
|
216
|
+
("invariance", self.invariance),
|
|
217
|
+
("variance", self.variance),
|
|
218
|
+
("covariance", self.covariance),
|
|
219
|
+
]:
|
|
220
|
+
if val < 0:
|
|
221
|
+
raise ValueError(f"{name} must be >= 0, got {val}")
|
|
222
|
+
|
|
223
|
+
@classmethod
|
|
224
|
+
def from_dict(cls, d: dict) -> VICRegConfig:
|
|
225
|
+
return cls(
|
|
226
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
227
|
+
augment=[AugmentConfig.from_dict(a) for a in d.get("augment", [])],
|
|
228
|
+
proj_dim=d.get("proj_dim", 256),
|
|
229
|
+
invariance=d.get("invariance", 25.0),
|
|
230
|
+
variance=d.get("variance", 25.0),
|
|
231
|
+
covariance=d.get("covariance", 1.0),
|
|
232
|
+
)
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
@dataclass
|
|
236
|
+
class BarlowTwinsConfig:
|
|
237
|
+
"""Validated configuration for Barlow Twins."""
|
|
238
|
+
|
|
239
|
+
encoder: EncoderConfig
|
|
240
|
+
augment: List[AugmentConfig] = field(default_factory=list)
|
|
241
|
+
proj_dim: int = 256
|
|
242
|
+
lambda_param: Optional[float] = None # defaults to 1/proj_dim inside the model
|
|
243
|
+
|
|
244
|
+
def __post_init__(self):
|
|
245
|
+
if self.proj_dim <= 0:
|
|
246
|
+
raise ValueError(f"proj_dim must be > 0, got {self.proj_dim}")
|
|
247
|
+
if self.lambda_param is not None and self.lambda_param < 0:
|
|
248
|
+
raise ValueError(f"lambda_param must be >= 0, got {self.lambda_param}")
|
|
249
|
+
|
|
250
|
+
@classmethod
|
|
251
|
+
def from_dict(cls, d: dict) -> BarlowTwinsConfig:
|
|
252
|
+
return cls(
|
|
253
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
254
|
+
augment=[AugmentConfig.from_dict(a) for a in d.get("augment", [])],
|
|
255
|
+
proj_dim=d.get("proj_dim", 256),
|
|
256
|
+
lambda_param=d.get("lambda_param", None),
|
|
257
|
+
)
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
@dataclass
|
|
261
|
+
class BGRLConfig:
|
|
262
|
+
"""Validated configuration for BGRL."""
|
|
263
|
+
|
|
264
|
+
encoder: EncoderConfig
|
|
265
|
+
augment: List[AugmentConfig] = field(default_factory=list)
|
|
266
|
+
pred_hidden: int = 512
|
|
267
|
+
ema_tau: float = 0.99
|
|
268
|
+
ema_tau_end: float = 1.0
|
|
269
|
+
total_steps: int = 0
|
|
270
|
+
|
|
271
|
+
def __post_init__(self):
|
|
272
|
+
if self.pred_hidden <= 0:
|
|
273
|
+
raise ValueError(f"pred_hidden must be > 0, got {self.pred_hidden}")
|
|
274
|
+
if not (0.0 < self.ema_tau < 1.0):
|
|
275
|
+
raise ValueError(f"ema_tau must be in (0, 1), got {self.ema_tau}")
|
|
276
|
+
if not (0.0 < self.ema_tau_end <= 1.0):
|
|
277
|
+
raise ValueError(f"ema_tau_end must be in (0, 1], got {self.ema_tau_end}")
|
|
278
|
+
if self.ema_tau > self.ema_tau_end:
|
|
279
|
+
raise ValueError(
|
|
280
|
+
f"ema_tau ({self.ema_tau}) must be <= ema_tau_end ({self.ema_tau_end})"
|
|
281
|
+
)
|
|
282
|
+
if self.total_steps < 0:
|
|
283
|
+
raise ValueError(f"total_steps must be >= 0, got {self.total_steps}")
|
|
284
|
+
|
|
285
|
+
@classmethod
|
|
286
|
+
def from_dict(cls, d: dict) -> BGRLConfig:
|
|
287
|
+
return cls(
|
|
288
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
289
|
+
augment=[AugmentConfig.from_dict(a) for a in d.get("augment", [])],
|
|
290
|
+
pred_hidden=d.get("pred_hidden", 512),
|
|
291
|
+
ema_tau=d.get("ema_tau", 0.99),
|
|
292
|
+
ema_tau_end=d.get("ema_tau_end", 1.0),
|
|
293
|
+
total_steps=d.get("total_steps", 0),
|
|
294
|
+
)
|
|
295
|
+
|
|
296
|
+
|
|
297
|
+
@dataclass
|
|
298
|
+
class AFGRLConfig:
|
|
299
|
+
"""Validated configuration for AFGRL."""
|
|
300
|
+
|
|
301
|
+
encoder: EncoderConfig
|
|
302
|
+
pred_hidden: int = 512
|
|
303
|
+
ema_tau: float = 0.99
|
|
304
|
+
ema_tau_end: float = 1.0
|
|
305
|
+
total_steps: int = 0
|
|
306
|
+
topk: int = 5
|
|
307
|
+
num_centroids: int = 50
|
|
308
|
+
num_kmeans: int = 4
|
|
309
|
+
clus_num_iters: int = 20
|
|
310
|
+
|
|
311
|
+
def __post_init__(self):
|
|
312
|
+
if self.pred_hidden <= 0:
|
|
313
|
+
raise ValueError(f"pred_hidden must be > 0, got {self.pred_hidden}")
|
|
314
|
+
if not (0.0 < self.ema_tau < 1.0):
|
|
315
|
+
raise ValueError(f"ema_tau must be in (0, 1), got {self.ema_tau}")
|
|
316
|
+
if not (0.0 < self.ema_tau_end <= 1.0):
|
|
317
|
+
raise ValueError(f"ema_tau_end must be in (0, 1], got {self.ema_tau_end}")
|
|
318
|
+
if self.ema_tau > self.ema_tau_end:
|
|
319
|
+
raise ValueError(
|
|
320
|
+
f"ema_tau ({self.ema_tau}) must be <= ema_tau_end ({self.ema_tau_end})"
|
|
321
|
+
)
|
|
322
|
+
if self.total_steps < 0:
|
|
323
|
+
raise ValueError(f"total_steps must be >= 0, got {self.total_steps}")
|
|
324
|
+
if self.topk <= 0:
|
|
325
|
+
raise ValueError(f"topk must be > 0, got {self.topk}")
|
|
326
|
+
if self.num_centroids <= 0:
|
|
327
|
+
raise ValueError(f"num_centroids must be > 0, got {self.num_centroids}")
|
|
328
|
+
if self.num_kmeans <= 0:
|
|
329
|
+
raise ValueError(f"num_kmeans must be > 0, got {self.num_kmeans}")
|
|
330
|
+
if self.clus_num_iters <= 0:
|
|
331
|
+
raise ValueError(f"clus_num_iters must be > 0, got {self.clus_num_iters}")
|
|
332
|
+
|
|
333
|
+
@classmethod
|
|
334
|
+
def from_dict(cls, d: dict) -> AFGRLConfig:
|
|
335
|
+
return cls(
|
|
336
|
+
encoder=EncoderConfig(**d["encoder"]),
|
|
337
|
+
pred_hidden=d.get("pred_hidden", 512),
|
|
338
|
+
ema_tau=d.get("ema_tau", 0.99),
|
|
339
|
+
ema_tau_end=d.get("ema_tau_end", 1.0),
|
|
340
|
+
total_steps=d.get("total_steps", 0),
|
|
341
|
+
topk=d.get("topk", 5),
|
|
342
|
+
num_centroids=d.get("num_centroids", 50),
|
|
343
|
+
num_kmeans=d.get("num_kmeans", 4),
|
|
344
|
+
clus_num_iters=d.get("clus_num_iters", 20),
|
|
345
|
+
)
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
@dataclass
|
|
349
|
+
class SupervisedConfig:
|
|
350
|
+
"""Validated configuration for the supervised baseline."""
|
|
351
|
+
|
|
352
|
+
encoder: EncoderConfig
|
|
353
|
+
|
|
354
|
+
@classmethod
|
|
355
|
+
def from_dict(cls, d: dict) -> SupervisedConfig:
|
|
356
|
+
return cls(encoder=EncoderConfig(**d["encoder"]))
|
|
357
|
+
|
|
358
|
+
|
|
359
|
+
@dataclass
|
|
360
|
+
class GraphDINOConfig:
|
|
361
|
+
"""Validated top-level configuration for GraphDINO."""
|
|
362
|
+
|
|
363
|
+
encoder: EncoderConfig
|
|
364
|
+
head: HeadConfig
|
|
365
|
+
augment_teacher: List[AugmentConfig] = field(default_factory=list)
|
|
366
|
+
augment_student: List[AugmentConfig] = field(default_factory=list)
|
|
367
|
+
ema_tau: float = 0.996
|
|
368
|
+
ema_tau_base: float = 0.996
|
|
369
|
+
total_steps: int = 0
|
|
370
|
+
freeze_last_layer_epochs: int = 1
|
|
371
|
+
n_views: int = 2
|
|
372
|
+
n_global_views: int = 2
|
|
373
|
+
|
|
374
|
+
def __post_init__(self):
|
|
375
|
+
if not (0.0 < self.ema_tau <= 1.0):
|
|
376
|
+
raise ValueError(f"ema_tau must be in (0, 1], got {self.ema_tau}")
|
|
377
|
+
if not (0.0 < self.ema_tau_base <= 1.0):
|
|
378
|
+
raise ValueError(f"ema_tau_base must be in (0, 1], got {self.ema_tau_base}")
|
|
379
|
+
if self.ema_tau_base > self.ema_tau:
|
|
380
|
+
raise ValueError(
|
|
381
|
+
f"ema_tau_base ({self.ema_tau_base}) must be <= ema_tau ({self.ema_tau})"
|
|
382
|
+
)
|
|
383
|
+
if self.total_steps < 0:
|
|
384
|
+
raise ValueError(f"total_steps must be >= 0, got {self.total_steps}")
|
|
385
|
+
if self.freeze_last_layer_epochs < 0:
|
|
386
|
+
raise ValueError(
|
|
387
|
+
f"freeze_last_layer_epochs must be >= 0, got {self.freeze_last_layer_epochs}"
|
|
388
|
+
)
|
|
389
|
+
if self.n_global_views < 1:
|
|
390
|
+
raise ValueError(f"n_global_views must be >= 1, got {self.n_global_views}")
|
|
391
|
+
if self.n_views < self.n_global_views:
|
|
392
|
+
raise ValueError(
|
|
393
|
+
f"n_views ({self.n_views}) must be >= n_global_views ({self.n_global_views})"
|
|
394
|
+
)
|
|
395
|
+
|
|
396
|
+
@classmethod
|
|
397
|
+
def from_dict(cls, d: dict) -> GraphDINOConfig:
|
|
398
|
+
encoder = EncoderConfig(**d["encoder"])
|
|
399
|
+
head = HeadConfig(**d["head"])
|
|
400
|
+
augment_teacher = [AugmentConfig.from_dict(a) for a in d.get("augment_teacher", [])]
|
|
401
|
+
augment_student = [AugmentConfig.from_dict(a) for a in d.get("augment_student", [])]
|
|
402
|
+
return cls(
|
|
403
|
+
encoder=encoder,
|
|
404
|
+
head=head,
|
|
405
|
+
augment_teacher=augment_teacher,
|
|
406
|
+
augment_student=augment_student,
|
|
407
|
+
ema_tau=d.get("ema_tau", 0.996),
|
|
408
|
+
ema_tau_base=d.get("ema_tau_base", d.get("ema_tau", 0.996)),
|
|
409
|
+
total_steps=d.get("total_steps", 0),
|
|
410
|
+
freeze_last_layer_epochs=d.get("freeze_last_layer_epochs", 1),
|
|
411
|
+
n_views=d.get("n_views", 2),
|
|
412
|
+
n_global_views=d.get("n_global_views", 2),
|
|
413
|
+
)
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from .augmentation import BaseAugmentation
|
|
2
|
+
from .callback import Callback
|
|
3
|
+
from .encoder import BaseEncoder
|
|
4
|
+
from .model import BaseModel, BaseSSLModel
|
|
5
|
+
from .registry import Registry
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"BaseAugmentation",
|
|
9
|
+
"BaseEncoder",
|
|
10
|
+
"BaseModel",
|
|
11
|
+
"BaseSSLModel",
|
|
12
|
+
"Callback",
|
|
13
|
+
"Registry",
|
|
14
|
+
]
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Protocol, runtime_checkable
|
|
4
|
+
|
|
5
|
+
from torch_geometric.data import Data
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@runtime_checkable
|
|
9
|
+
class BaseAugmentation(Protocol):
|
|
10
|
+
"""Structural interface for graph augmentation transforms.
|
|
11
|
+
|
|
12
|
+
Any callable Data -> Data satisfies this protocol, matching the
|
|
13
|
+
torchvision-style API used in augmentation/transforms.py.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
def __call__(self, data: Data) -> Data: ...
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import TYPE_CHECKING, Any, Dict
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
if TYPE_CHECKING:
|
|
8
|
+
from torch_geometric.data import Data
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class Callback:
|
|
12
|
+
"""Base class for Trainer callbacks with no-op defaults.
|
|
13
|
+
|
|
14
|
+
Subclass and override only the hooks you need.
|
|
15
|
+
All methods receive the trainer instance as first argument so callbacks
|
|
16
|
+
can access trainer.device, trainer.optimizer, etc. without coupling.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
def on_train_start(self, trainer: Any, model: torch.nn.Module) -> None:
|
|
20
|
+
pass
|
|
21
|
+
|
|
22
|
+
def on_epoch_start(self, trainer: Any, model: torch.nn.Module, epoch: int) -> None:
|
|
23
|
+
pass
|
|
24
|
+
|
|
25
|
+
def on_batch_end(
|
|
26
|
+
self,
|
|
27
|
+
trainer: Any,
|
|
28
|
+
model: torch.nn.Module,
|
|
29
|
+
loss: torch.Tensor,
|
|
30
|
+
batch: "Data",
|
|
31
|
+
) -> None:
|
|
32
|
+
pass
|
|
33
|
+
|
|
34
|
+
def on_epoch_end(
|
|
35
|
+
self,
|
|
36
|
+
trainer: Any,
|
|
37
|
+
model: torch.nn.Module,
|
|
38
|
+
epoch: int,
|
|
39
|
+
metrics: Dict[str, Any],
|
|
40
|
+
) -> None:
|
|
41
|
+
pass
|
|
42
|
+
|
|
43
|
+
def on_train_end(self, trainer: Any, model: torch.nn.Module) -> None:
|
|
44
|
+
pass
|
graphssl/core/encoder.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Optional, Protocol, runtime_checkable
|
|
4
|
+
|
|
5
|
+
from torch import Tensor
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@runtime_checkable
|
|
9
|
+
class BaseEncoder(Protocol):
|
|
10
|
+
"""Structural interface for GNN backbone encoders.
|
|
11
|
+
|
|
12
|
+
Using Protocol (not ABC) so existing nn.Module encoders conform without
|
|
13
|
+
inheriting from this class — structural subtyping only.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
def forward(self, x: Tensor, edge_index: Tensor, batch: Optional[Tensor] = None) -> Tensor: ...
|
|
17
|
+
|
|
18
|
+
def reset_parameters(self) -> None: ...
|
graphssl/core/model.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from abc import ABC, abstractmethod
|
|
4
|
+
from typing import Iterator
|
|
5
|
+
|
|
6
|
+
import torch.nn as nn
|
|
7
|
+
from torch import Tensor
|
|
8
|
+
from torch_geometric.data import Data
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class BaseModel(nn.Module, ABC):
|
|
12
|
+
"""Base for all models. forward() returns embeddings only."""
|
|
13
|
+
|
|
14
|
+
@abstractmethod
|
|
15
|
+
def forward(self, data: Data) -> Tensor:
|
|
16
|
+
"""Extract node/graph embeddings from a batch."""
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class BaseSSLModel(BaseModel):
|
|
20
|
+
"""Extension for self-supervised models.
|
|
21
|
+
|
|
22
|
+
The Trainer loop is:
|
|
23
|
+
model.on_epoch_start(epoch) # teacher temp warmup, etc.
|
|
24
|
+
for batch in loader:
|
|
25
|
+
loss = model.compute_loss(batch)
|
|
26
|
+
loss.backward()
|
|
27
|
+
clip_grad_norm_(...)
|
|
28
|
+
model.post_backward() # freeze last layer, etc.
|
|
29
|
+
optimizer.step()
|
|
30
|
+
model.post_step() # EMA, center updates, etc.
|
|
31
|
+
model.on_epoch_end(epoch) # epoch counter, etc.
|
|
32
|
+
|
|
33
|
+
For evaluation:
|
|
34
|
+
embeddings = model(batch) # calls forward(), no grad needed
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
@abstractmethod
|
|
38
|
+
def compute_loss(self, data: Data) -> Tensor:
|
|
39
|
+
"""Compute the SSL training loss. Augmentation is handled internally."""
|
|
40
|
+
|
|
41
|
+
@abstractmethod
|
|
42
|
+
def student_parameters(self) -> Iterator[nn.Parameter]:
|
|
43
|
+
"""Parameters to optimize. Excludes frozen teacher/target parameters."""
|
|
44
|
+
|
|
45
|
+
def post_backward(self) -> None:
|
|
46
|
+
"""Called after backward() and grad clipping, before optimizer.step().
|
|
47
|
+
|
|
48
|
+
Override to cancel gradients on specific parameters (e.g. DINO last layer freeze).
|
|
49
|
+
"""
|
|
50
|
+
|
|
51
|
+
def post_step(self) -> None:
|
|
52
|
+
"""Called after optimizer.step(). Override for EMA updates, center updates, etc."""
|
|
53
|
+
|
|
54
|
+
def on_epoch_start(self, epoch: int) -> None:
|
|
55
|
+
"""Called by Trainer at the start of each epoch before any batch."""
|
|
56
|
+
|
|
57
|
+
def on_epoch_end(self, epoch: int) -> None:
|
|
58
|
+
"""Called by Trainer at the end of each epoch after all batches."""
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Callable, Dict, Generic, Iterator, TypeVar
|
|
4
|
+
|
|
5
|
+
T = TypeVar("T")
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Registry(Generic[T]):
|
|
9
|
+
"""Generic registry mapping string names to builder callables.
|
|
10
|
+
|
|
11
|
+
Usage:
|
|
12
|
+
ENCODERS: Registry = Registry()
|
|
13
|
+
|
|
14
|
+
@ENCODERS.register("gcn")
|
|
15
|
+
class GCNEncoder(nn.Module): ...
|
|
16
|
+
|
|
17
|
+
enc = ENCODERS.build("gcn", in_channels=32, hidden_dim=64)
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
def __init__(self) -> None:
|
|
21
|
+
self._builders: Dict[str, Callable[..., T]] = {}
|
|
22
|
+
|
|
23
|
+
def register(self, name: str) -> Callable[[Callable[..., T]], Callable[..., T]]:
|
|
24
|
+
def deco(fn: Callable[..., T]) -> Callable[..., T]:
|
|
25
|
+
if name in self._builders:
|
|
26
|
+
raise KeyError(f"Duplicate registration: {name!r}")
|
|
27
|
+
self._builders[name] = fn
|
|
28
|
+
return fn
|
|
29
|
+
|
|
30
|
+
return deco
|
|
31
|
+
|
|
32
|
+
def get_builder(self, name: str) -> Callable[..., T]:
|
|
33
|
+
if name not in self._builders:
|
|
34
|
+
raise KeyError(f"Unknown component {name!r}. Available: {list(self._builders)}")
|
|
35
|
+
return self._builders[name]
|
|
36
|
+
|
|
37
|
+
def build(self, name: str, **kwargs) -> T:
|
|
38
|
+
return self.get_builder(name)(**kwargs)
|
|
39
|
+
|
|
40
|
+
def __contains__(self, name: str) -> bool:
|
|
41
|
+
return name in self._builders
|
|
42
|
+
|
|
43
|
+
def __iter__(self) -> Iterator[str]:
|
|
44
|
+
return iter(self._builders)
|
|
45
|
+
|
|
46
|
+
def __repr__(self) -> str:
|
|
47
|
+
return f"Registry({list(self._builders)})"
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from .datamodule import DataModule
|