D4CMPP2 0.4.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 (103) hide show
  1. D4CMPP2/_Data/AGENTS.md +24 -0
  2. D4CMPP2/_Data/Aqsoldb.csv +9291 -0
  3. D4CMPP2/_Data/BradleyMP.csv +3042 -0
  4. D4CMPP2/_Data/Lipophilicity.csv +1131 -0
  5. D4CMPP2/_Data/README.md +8 -0
  6. D4CMPP2/_Data/__init__.py +26 -0
  7. D4CMPP2/_Data/optical.csv +20237 -0
  8. D4CMPP2/_Data/test.csv +190 -0
  9. D4CMPP2/__init__.py +16 -0
  10. D4CMPP2/__main__.py +5 -0
  11. D4CMPP2/_main.py +500 -0
  12. D4CMPP2/cli.py +7 -0
  13. D4CMPP2/exceptions.py +53 -0
  14. D4CMPP2/grid_search.py +259 -0
  15. D4CMPP2/network_refer.yaml +160 -0
  16. D4CMPP2/networks/AFP_model.py +72 -0
  17. D4CMPP2/networks/AFPwithSolv_model.py +72 -0
  18. D4CMPP2/networks/DMPNN_model.py +90 -0
  19. D4CMPP2/networks/DMPNNwithSolv_model.py +89 -0
  20. D4CMPP2/networks/GAT_model.py +48 -0
  21. D4CMPP2/networks/GATwithSolv_model.py +63 -0
  22. D4CMPP2/networks/GCN_model.py +113 -0
  23. D4CMPP2/networks/GCNwithSolv_model.py +103 -0
  24. D4CMPP2/networks/GC_model.py +122 -0
  25. D4CMPP2/networks/ISATPM_model.py +14 -0
  26. D4CMPP2/networks/ISATPN_model.py +199 -0
  27. D4CMPP2/networks/ISAT_model.py +90 -0
  28. D4CMPP2/networks/MPNN_model.py +56 -0
  29. D4CMPP2/networks/MPNNwithSolv_model.py +72 -0
  30. D4CMPP2/networks/__init__.py +25 -0
  31. D4CMPP2/networks/base.py +250 -0
  32. D4CMPP2/networks/registry.py +187 -0
  33. D4CMPP2/networks/src/AFP.py +118 -0
  34. D4CMPP2/networks/src/BiDropout.py +29 -0
  35. D4CMPP2/networks/src/DMPNN.py +35 -0
  36. D4CMPP2/networks/src/GAT.py +69 -0
  37. D4CMPP2/networks/src/GC.py +85 -0
  38. D4CMPP2/networks/src/GCN.py +71 -0
  39. D4CMPP2/networks/src/ISAT.py +153 -0
  40. D4CMPP2/networks/src/Linear.py +49 -0
  41. D4CMPP2/networks/src/MPNN.py +56 -0
  42. D4CMPP2/networks/src/SolventLayer.py +62 -0
  43. D4CMPP2/networks/src/__init__.py +0 -0
  44. D4CMPP2/networks/src/distGCN.py +21 -0
  45. D4CMPP2/networks/src/pyg_hetero.py +24 -0
  46. D4CMPP2/optimize.py +472 -0
  47. D4CMPP2/src/Analyzer/ISAAnalyzer.py +458 -0
  48. D4CMPP2/src/Analyzer/ISAPNAnalyzer.py +366 -0
  49. D4CMPP2/src/Analyzer/ISAwSAnalyzer.py +117 -0
  50. D4CMPP2/src/Analyzer/MolAnalyzer.py +319 -0
  51. D4CMPP2/src/Analyzer/__init__.py +54 -0
  52. D4CMPP2/src/Analyzer/core.py +480 -0
  53. D4CMPP2/src/Analyzer/factory.py +166 -0
  54. D4CMPP2/src/Analyzer/interpretation.py +232 -0
  55. D4CMPP2/src/Analyzer/results.py +101 -0
  56. D4CMPP2/src/DataManager/Dataset/GraphDataset.py +314 -0
  57. D4CMPP2/src/DataManager/Dataset/ISAGraphDataset.py +384 -0
  58. D4CMPP2/src/DataManager/Dataset/__init__.py +0 -0
  59. D4CMPP2/src/DataManager/GraphGenerator/ISAGraphGenerator.py +223 -0
  60. D4CMPP2/src/DataManager/GraphGenerator/MolGraphGenerator.py +73 -0
  61. D4CMPP2/src/DataManager/GraphGenerator/__init__.py +14 -0
  62. D4CMPP2/src/DataManager/ISADataManager.py +67 -0
  63. D4CMPP2/src/DataManager/MolDataManager.py +735 -0
  64. D4CMPP2/src/DataManager/__init__.py +14 -0
  65. D4CMPP2/src/DataManager/contracts.py +179 -0
  66. D4CMPP2/src/NetworkManager/ISANetworkManager.py +12 -0
  67. D4CMPP2/src/NetworkManager/NetworkManager.py +520 -0
  68. D4CMPP2/src/NetworkManager/__init__.py +14 -0
  69. D4CMPP2/src/PostProcessor.py +160 -0
  70. D4CMPP2/src/TrainManager/ISATrainManager.py +26 -0
  71. D4CMPP2/src/TrainManager/TrainManager.py +254 -0
  72. D4CMPP2/src/TrainManager/__init__.py +14 -0
  73. D4CMPP2/src/TrainManager/callbacks.py +119 -0
  74. D4CMPP2/src/__init__.py +0 -0
  75. D4CMPP2/src/utils/PATH.py +246 -0
  76. D4CMPP2/src/utils/__init__.py +0 -0
  77. D4CMPP2/src/utils/argparser.py +56 -0
  78. D4CMPP2/src/utils/checkpointing.py +90 -0
  79. D4CMPP2/src/utils/config_resolution.py +123 -0
  80. D4CMPP2/src/utils/config_validation.py +370 -0
  81. D4CMPP2/src/utils/csv_validation.py +105 -0
  82. D4CMPP2/src/utils/data_quality.py +181 -0
  83. D4CMPP2/src/utils/featureizer.py +202 -0
  84. D4CMPP2/src/utils/functional_group.csv +169 -0
  85. D4CMPP2/src/utils/graph_cache.py +213 -0
  86. D4CMPP2/src/utils/leaderboard.py +212 -0
  87. D4CMPP2/src/utils/metrics.py +31 -0
  88. D4CMPP2/src/utils/module_loader.py +147 -0
  89. D4CMPP2/src/utils/output.py +80 -0
  90. D4CMPP2/src/utils/reproducibility.py +70 -0
  91. D4CMPP2/src/utils/run_manifest.py +175 -0
  92. D4CMPP2/src/utils/scaler.py +40 -0
  93. D4CMPP2/src/utils/sculptor.py +713 -0
  94. D4CMPP2/src/utils/splitting.py +250 -0
  95. D4CMPP2/src/utils/supportfile_saver.py +94 -0
  96. D4CMPP2/src/utils/tools.py +156 -0
  97. D4CMPP2/src/utils/transfer_learning.py +111 -0
  98. d4cmpp2-0.4.0.dist-info/METADATA +420 -0
  99. d4cmpp2-0.4.0.dist-info/RECORD +103 -0
  100. d4cmpp2-0.4.0.dist-info/WHEEL +5 -0
  101. d4cmpp2-0.4.0.dist-info/entry_points.txt +2 -0
  102. d4cmpp2-0.4.0.dist-info/licenses/LICENSE +21 -0
  103. d4cmpp2-0.4.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,175 @@
1
+ """Additive, best-effort run manifest recording."""
2
+
3
+ import hashlib
4
+ import json
5
+ import os
6
+ import platform
7
+ import subprocess
8
+ import sys
9
+ import time
10
+ import uuid
11
+ import warnings
12
+ from datetime import datetime, timezone
13
+ from pathlib import Path
14
+
15
+
16
+ MANIFEST_SCHEMA_VERSION = 1
17
+
18
+
19
+ def _json_safe(value):
20
+ if isinstance(value, (str, int, float, bool)) or value is None:
21
+ return value
22
+ if isinstance(value, os.PathLike):
23
+ return os.fspath(value)
24
+ if isinstance(value, dict):
25
+ result = {}
26
+ for key, item in value.items():
27
+ name = str(key)
28
+ if any(marker in name.lower() for marker in ("password", "token", "secret", "credential")):
29
+ result[name] = "<redacted>"
30
+ else:
31
+ result[name] = _json_safe(item)
32
+ return result
33
+ if isinstance(value, (list, tuple)):
34
+ return [_json_safe(item) for item in value]
35
+ return repr(value)
36
+
37
+
38
+ def _version(module_name):
39
+ try:
40
+ module = __import__(module_name)
41
+ return str(getattr(module, "__version__", "unknown"))
42
+ except (ImportError, OSError):
43
+ return None
44
+
45
+
46
+ def _file_hash(path):
47
+ if not path or not Path(path).is_file():
48
+ return None
49
+ digest = hashlib.sha256()
50
+ with open(path, "rb") as stream:
51
+ for chunk in iter(lambda: stream.read(1024 * 1024), b""):
52
+ digest.update(chunk)
53
+ return digest.hexdigest()
54
+
55
+
56
+ def _git_state(root):
57
+ try:
58
+ commit = subprocess.run(
59
+ ["git", "-c", f"safe.directory={root}", "-C", str(root), "rev-parse", "HEAD"],
60
+ capture_output=True,
61
+ text=True,
62
+ check=True,
63
+ timeout=5,
64
+ ).stdout.strip()
65
+ dirty = bool(
66
+ subprocess.run(
67
+ ["git", "-c", f"safe.directory={root}", "-C", str(root), "status", "--porcelain"],
68
+ capture_output=True,
69
+ text=True,
70
+ check=True,
71
+ timeout=5,
72
+ ).stdout.strip()
73
+ )
74
+ return {"commit": commit, "dirty": dirty}
75
+ except (OSError, subprocess.SubprocessError):
76
+ return {"commit": None, "dirty": None}
77
+
78
+
79
+ class RunManifest:
80
+ def __init__(self, config, mode):
81
+ self.started = time.time()
82
+ stamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S")
83
+ self.run_id = f"{stamp}-{uuid.uuid4().hex[:12]}"
84
+ self.path = Path(config["MODEL_PATH"]) / "runs" / self.run_id / "run_manifest.json"
85
+ package_root = Path(__file__).resolve().parents[2]
86
+ self.data = {
87
+ "manifest_schema_version": MANIFEST_SCHEMA_VERSION,
88
+ "run_id": self.run_id,
89
+ "status": "running",
90
+ "mode": mode,
91
+ "started_at": datetime.now(timezone.utc).isoformat(),
92
+ "config": _json_safe(config),
93
+ "network": {
94
+ "id": config.get("network_id", config.get("network")),
95
+ "module": config.get("network"),
96
+ "data_manager": config.get("data_manager_class"),
97
+ "network_manager": config.get("network_manager_class"),
98
+ "train_manager": config.get("train_manager_class"),
99
+ },
100
+ "data": {
101
+ "path": config.get("DATA_PATH"),
102
+ "sha256": _file_hash(config.get("DATA_PATH")),
103
+ "targets": _json_safe(config.get("target")),
104
+ },
105
+ "graph": {
106
+ "backend": "pyg",
107
+ "schema_version": 2,
108
+ "cache_directory": config.get("GRAPH_DIR"),
109
+ "cache_policy": config.get("graph_cache_policy", "v2"),
110
+ },
111
+ "seeds": {
112
+ "split_random_seed": config.get("split_random_seed", 42),
113
+ "random_seed": config.get("random_seed"),
114
+ "scheduler_policy": config.get("scheduler_policy", "legacy_dual"),
115
+ "effective_reproducibility": _json_safe(
116
+ config.get("effective_reproducibility", {})
117
+ ),
118
+ },
119
+ "split": {
120
+ "strategy": config.get("split_strategy", "auto"),
121
+ "scaffold_column": config.get("scaffold_column"),
122
+ "scaffold_include_chirality": config.get(
123
+ "scaffold_include_chirality", False
124
+ ),
125
+ },
126
+ "environment": {
127
+ "python": sys.version.split()[0],
128
+ "torch": _version("torch"),
129
+ "torch_geometric": _version("torch_geometric"),
130
+ "rdkit": _version("rdkit"),
131
+ "numpy": _version("numpy"),
132
+ "pandas": _version("pandas"),
133
+ "os": platform.platform(),
134
+ "device": config.get("device"),
135
+ "git": _git_state(package_root),
136
+ },
137
+ }
138
+ self.write()
139
+
140
+ def update(self, **values):
141
+ self.data.update(_json_safe(values))
142
+ self.write()
143
+
144
+ def finish(self, status, error=None, **values):
145
+ self.data.update(_json_safe(values))
146
+ self.data["status"] = status
147
+ self.data["ended_at"] = datetime.now(timezone.utc).isoformat()
148
+ self.data["duration_seconds"] = time.time() - self.started
149
+ if error is not None:
150
+ self.data["error"] = {
151
+ "type": type(error).__name__,
152
+ "message": str(error)[:2000],
153
+ }
154
+ self.write()
155
+
156
+ def write(self):
157
+ try:
158
+ self.path.parent.mkdir(parents=True, exist_ok=True)
159
+ staging = self.path.with_name(f".{self.path.name}.{uuid.uuid4().hex}.tmp")
160
+ try:
161
+ staging.write_text(
162
+ json.dumps(self.data, indent=2, sort_keys=True),
163
+ encoding="utf-8",
164
+ )
165
+ os.replace(staging, self.path)
166
+ finally:
167
+ if staging.exists():
168
+ staging.unlink()
169
+ except OSError as exc:
170
+ warnings.warn(
171
+ f"Run manifest {str(self.path)!r} could not be written: {exc}. "
172
+ "Training/checkpoint results are unaffected.",
173
+ RuntimeWarning,
174
+ stacklevel=2,
175
+ )
@@ -0,0 +1,40 @@
1
+ from sklearn.preprocessing import StandardScaler, MinMaxScaler, Normalizer, RobustScaler
2
+
3
+ class identityScaler:
4
+ def fit(self,X):
5
+ pass
6
+ def transform(self,X):
7
+ return X
8
+ def fit_transform(self,X):
9
+ return X
10
+ def inverse_transform(self,X):
11
+ return X
12
+
13
+
14
+ class Scaler:
15
+ def __init__(self,scale_type="standard"):
16
+ self.scale_type = scale_type
17
+ if self.scale_type=='standard':
18
+ self.scaler = StandardScaler()
19
+ elif self.scale_type=='minmax':
20
+ self.scaler = MinMaxScaler()
21
+ elif self.scale_type=='normalizer':
22
+ self.scaler = Normalizer()
23
+ elif self.scale_type=='robust':
24
+ self.scaler = RobustScaler()
25
+ else:
26
+ self.scaler = identityScaler()
27
+
28
+ def fit(self,X):
29
+ self.scaler.fit(X)
30
+
31
+ def transform(self,X):
32
+ return self.scaler.transform(X)
33
+
34
+ def fit_transform(self,X):
35
+ return self.scaler.fit_transform(X)
36
+
37
+ def inverse_transform(self,X):
38
+ return self.scaler.inverse_transform(X)
39
+
40
+