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.
Files changed (59) hide show
  1. graphssl/__init__.py +63 -0
  2. graphssl/augmentation/__init__.py +11 -0
  3. graphssl/augmentation/compose.py +60 -0
  4. graphssl/augmentation/functional.py +127 -0
  5. graphssl/augmentation/transforms.py +84 -0
  6. graphssl/config/__init__.py +14 -0
  7. graphssl/config/load.py +114 -0
  8. graphssl/config/schema.py +413 -0
  9. graphssl/core/__init__.py +14 -0
  10. graphssl/core/augmentation.py +16 -0
  11. graphssl/core/callback.py +44 -0
  12. graphssl/core/encoder.py +18 -0
  13. graphssl/core/model.py +58 -0
  14. graphssl/core/registry.py +47 -0
  15. graphssl/data/__init__.py +1 -0
  16. graphssl/data/datamodule.py +100 -0
  17. graphssl/encoders/__init__.py +3 -0
  18. graphssl/encoders/gcn.py +104 -0
  19. graphssl/encoders/gin.py +181 -0
  20. graphssl/encoders/transformer.py +206 -0
  21. graphssl/evaluation/__init__.py +3 -0
  22. graphssl/evaluation/knn.py +52 -0
  23. graphssl/evaluation/linear_probe.py +116 -0
  24. graphssl/evaluation/visualization.py +206 -0
  25. graphssl/losses/__init__.py +6 -0
  26. graphssl/losses/barlow.py +35 -0
  27. graphssl/losses/combined.py +67 -0
  28. graphssl/losses/dino.py +38 -0
  29. graphssl/losses/nt_xent.py +40 -0
  30. graphssl/losses/regression.py +44 -0
  31. graphssl/losses/vicreg.py +56 -0
  32. graphssl/models/__init__.py +14 -0
  33. graphssl/models/afgrl.py +133 -0
  34. graphssl/models/barlow_twins.py +68 -0
  35. graphssl/models/bgrl.py +110 -0
  36. graphssl/models/dgi.py +85 -0
  37. graphssl/models/graphcl.py +68 -0
  38. graphssl/models/graphdino.py +182 -0
  39. graphssl/models/supervised.py +48 -0
  40. graphssl/models/vicreg.py +72 -0
  41. graphssl/nn/__init__.py +4 -0
  42. graphssl/nn/dino_head.py +99 -0
  43. graphssl/nn/mlp.py +52 -0
  44. graphssl/nn/norm.py +24 -0
  45. graphssl/nn/pooling.py +18 -0
  46. graphssl/registry/__init__.py +1 -0
  47. graphssl/registry/registry.py +9 -0
  48. graphssl/training/__init__.py +2 -0
  49. graphssl/training/callbacks.py +138 -0
  50. graphssl/training/trainer.py +108 -0
  51. graphssl/utils/__init__.py +3 -0
  52. graphssl/utils/ema.py +14 -0
  53. graphssl/utils/positive_miner.py +119 -0
  54. graphssl/utils/schedulers.py +55 -0
  55. graphssl-0.1.0.dist-info/METADATA +273 -0
  56. graphssl-0.1.0.dist-info/RECORD +59 -0
  57. graphssl-0.1.0.dist-info/WHEEL +5 -0
  58. graphssl-0.1.0.dist-info/licenses/LICENSE +21 -0
  59. 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
@@ -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