jorwal 1.0.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.
__init__.py ADDED
File without changes
admet_critic.py ADDED
@@ -0,0 +1,122 @@
1
+ """
2
+ Agentic BioNeMo - Agent 3: ADMET & MedChem Critic Agent
3
+ Rigorous RDKit-based physicochemical property profiling, Lipinski/Veber rules,
4
+ PAINS structural alerts, and ADMET toxicity heuristics.
5
+ """
6
+ import logging
7
+ from typing import List, Tuple
8
+ from rdkit import Chem
9
+ from rdkit.Chem import Descriptors, Crippen, rdMolDescriptors, QED
10
+ from src.models import MoleculeCandidate, AgentMessage
11
+
12
+ logger = logging.getLogger("ADMETCritic")
13
+
14
+ # PAINS and reactive motif SMARTS patterns
15
+ PAINS_PATTERNS = {
16
+ "Quinone": "[#6]1(=[O,S])[#6]=[#6][#6](=[O,S])[#6]=[#6]1",
17
+ "Rhodanine": "O=C1NC(=S)SC1",
18
+ "Michael Acceptor": "[#6]=[#6]-[#6](=[O,S])",
19
+ "Aliphatic Halide": "[CX4][Cl,Br,I]",
20
+ "Epoxide": "C1OC1",
21
+ "Catechol": "c1c([OX2H])c([OX2H])ccc1"
22
+ }
23
+
24
+ class ADMETCriticAgent:
25
+ def __init__(self, name: str = "ADMETCritic"):
26
+ self.name = name
27
+ self._compiled_pains = {k: Chem.MolFromSmarts(v) for k, v in PAINS_PATTERNS.items() if Chem.MolFromSmarts(v)}
28
+
29
+ def evaluate_candidates(self, candidates: List[MoleculeCandidate]) -> Tuple[List[MoleculeCandidate], AgentMessage]:
30
+ """Runs full ADMET screening across all generated molecule candidates."""
31
+ passed_count = 0
32
+ flagged_count = 0
33
+ rejected_count = 0
34
+
35
+ for cand in candidates:
36
+ mol = Chem.MolFromSmiles(cand.smiles)
37
+ if not mol:
38
+ cand.admet_verdict = "REJECT"
39
+ cand.admet_notes.append("Invalid chemical structure / SMILES parse error.")
40
+ rejected_count += 1
41
+ continue
42
+
43
+ try:
44
+ Chem.SanitizeMol(mol)
45
+ except Exception as e:
46
+ cand.admet_verdict = "REJECT"
47
+ cand.admet_notes.append(f"Valence sanitization failed: {e}")
48
+ rejected_count += 1
49
+ continue
50
+
51
+ # Descriptors
52
+ cand.mw = round(Descriptors.MolWt(mol), 2)
53
+ cand.logp = round(Crippen.MolLogP(mol), 2)
54
+ cand.tpsa = round(rdMolDescriptors.CalcTPSA(mol), 2)
55
+ cand.hbd = rdMolDescriptors.CalcNumHBD(mol)
56
+ cand.hba = rdMolDescriptors.CalcNumHBA(mol)
57
+ cand.rotatable_bonds = rdMolDescriptors.CalcNumRotatableBonds(mol)
58
+ cand.qed = round(QED.qed(mol), 3)
59
+ cand.molecular_formula = rdMolDescriptors.CalcMolFormula(mol)
60
+
61
+ # Synthetic Accessibility Heuristic (Ring complexity + sp3 ratio + rotatable bonds)
62
+ fsp3 = rdMolDescriptors.CalcFractionCSP3(mol)
63
+ ring_count = rdMolDescriptors.CalcNumRings(mol)
64
+ # SAScore heuristic scale: 1 (easy) to 10 (hard)
65
+ sascore = 2.0 + (cand.mw / 150.0) + (ring_count * 0.4) - (fsp3 * 1.5)
66
+ cand.sascore = round(max(1.0, min(10.0, sascore)), 2)
67
+
68
+ # PAINS alerts
69
+ cand.pains_alerts = []
70
+ for p_name, p_mol in self._compiled_pains.items():
71
+ if mol.HasSubstructMatch(p_mol):
72
+ cand.pains_alerts.append(p_name)
73
+
74
+ # Lipinski & Veber
75
+ lipinski_violations = 0
76
+ if cand.mw > 550: lipinski_violations += 1
77
+ if cand.logp > 5.0: lipinski_violations += 1
78
+ if cand.hbd > 5: lipinski_violations += 1
79
+ if cand.hba > 10: lipinski_violations += 1
80
+ cand.passes_lipinski = (lipinski_violations <= 1)
81
+
82
+ cand.passes_veber = (cand.rotatable_bonds <= 10 and cand.tpsa <= 140.0)
83
+
84
+ # Blood-Brain Barrier (BBB) Heuristic
85
+ cand.bbb_permeable = (1.5 <= cand.logp <= 3.8 and cand.tpsa < 90.0 and cand.mw < 450)
86
+
87
+ # hERG Cardiotoxicity Alert Heuristic (High LogP + high MW + basic amines)
88
+ cand.herg_liability = (cand.logp > 4.2 and cand.mw > 480 and cand.hba >= 4)
89
+
90
+ # Final Verdict
91
+ cand.admet_notes = []
92
+ if cand.pains_alerts:
93
+ cand.admet_verdict = "REJECT"
94
+ cand.admet_notes.append(f"PAINS alerts: {', '.join(cand.pains_alerts)}")
95
+ rejected_count += 1
96
+ elif not cand.passes_lipinski or not cand.passes_veber:
97
+ cand.admet_verdict = "FLAGGED"
98
+ cand.admet_notes.append(f"Lipinski/Veber boundary: MW={cand.mw}, LogP={cand.logp}, TPSA={cand.tpsa}")
99
+ flagged_count += 1
100
+ elif cand.herg_liability:
101
+ cand.admet_verdict = "FLAGGED"
102
+ cand.admet_notes.append(f"hERG potential cardiotoxicity liability (LogP={cand.logp}, MW={cand.mw})")
103
+ flagged_count += 1
104
+ else:
105
+ cand.admet_verdict = "PASS"
106
+ cand.admet_notes.append(f"Compliant: QED={cand.qed}, SAScore={cand.sascore}, TPSA={cand.tpsa}")
107
+ passed_count += 1
108
+
109
+ thought_text = (
110
+ f"Evaluated {len(candidates)} candidates through medicinal chemistry filtration. "
111
+ f"Outcome: {passed_count} PASS, {flagged_count} FLAGGED, {rejected_count} REJECT. "
112
+ f"Screened PAINS, Lipinski Rule of 5, Veber criteria, SAScore, and hERG liability."
113
+ )
114
+ message = AgentMessage(
115
+ agent_name=self.name,
116
+ role="ADMET Critic",
117
+ action="ADMET_MEDCHEM_EVALUATION",
118
+ thought=thought_text,
119
+ output_summary=f"{passed_count} PASS + {flagged_count} FLAGGED leads cleared for molecular docking simulation.",
120
+ status="SUCCESS"
121
+ )
122
+ return candidates, message
baseline.py ADDED
@@ -0,0 +1,222 @@
1
+ """
2
+ Local CPU Baseline Benchmark Runner for Protein Embedding Generation.
3
+
4
+ Benchmarks Hugging Face's lightweight `facebook/esm2_t6_8M_UR50D` model on local CPU
5
+ without GPU acceleration to evaluate time-per-residue (ms/residue), token throughput
6
+ (tokens/sec), and memory footprint.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import time
12
+ import logging
13
+ import tracemalloc
14
+ from typing import Dict, Any, Optional, Tuple
15
+ from dataclasses import dataclass, field
16
+
17
+ import torch
18
+
19
+ logger = logging.getLogger(__name__)
20
+
21
+ DEFAULT_BASELINE_MODEL = "facebook/esm2_t6_8M_UR50D"
22
+
23
+
24
+ @dataclass
25
+ class BaselineProfileResult:
26
+ """Latency, throughput, and memory profiling metrics for CPU baseline inference."""
27
+ sequence_length: int
28
+ tokenization_ms: float
29
+ forward_pass_ms: float
30
+ total_cpu_latency_ms: float
31
+ ms_per_residue: float
32
+ tokens_per_sec: float
33
+ peak_memory_mb: float
34
+ tensor_bytes: int
35
+ tensor_shape: Tuple[int, ...]
36
+ model_name: str
37
+ is_mock: bool = False
38
+ metadata: Dict[str, Any] = field(default_factory=dict)
39
+
40
+
41
+ class LocalCPUBaseline:
42
+ """
43
+ Runner for local CPU ESM-2 inference benchmarking.
44
+ Measures baseline unaccelerated compute performance across varying protein sequence lengths.
45
+ """
46
+
47
+ def __init__(
48
+ self,
49
+ model_name: str = DEFAULT_BASELINE_MODEL,
50
+ device: str = "cpu",
51
+ lazy_load: bool = True,
52
+ mock_fallback: bool = True,
53
+ ):
54
+ self.model_name = model_name
55
+ self.device = torch.device(device)
56
+ self.mock_fallback = mock_fallback
57
+ self.tokenizer = None
58
+ self.model = None
59
+ self._is_loaded = False
60
+ self._using_mock = False
61
+
62
+ if not lazy_load:
63
+ self.load_model()
64
+
65
+ def load_model(self) -> None:
66
+ """Loads tokenizer and model weights onto the target device (CPU)."""
67
+ if self._is_loaded:
68
+ return
69
+
70
+ try:
71
+ logger.info("Loading baseline ESM-2 model '%s' on %s...", self.model_name, self.device)
72
+ from transformers import AutoTokenizer, AutoModel
73
+
74
+ self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)
75
+ self.model = AutoModel.from_pretrained(self.model_name)
76
+ self.model.to(self.device)
77
+ self.model.eval()
78
+ self._is_loaded = True
79
+ self._using_mock = False
80
+ logger.info("Baseline ESM-2 model successfully loaded.")
81
+ except Exception as e:
82
+ if self.mock_fallback:
83
+ logger.warning(
84
+ "Unable to load '%s' from HuggingFace hub (%s). Falling back to CPU simulation mode.",
85
+ self.model_name,
86
+ str(e),
87
+ )
88
+ self._is_loaded = True
89
+ self._using_mock = True
90
+ else:
91
+ raise RuntimeError(f"Failed to load baseline model {self.model_name}: {e}") from e
92
+
93
+ def benchmark_sequence(
94
+ self,
95
+ sequence: str,
96
+ warmup_runs: int = 1,
97
+ benchmark_runs: int = 3,
98
+ ) -> BaselineProfileResult:
99
+ """
100
+ Runs protein sequence embedding inference on local CPU and records execution metrics.
101
+
102
+ Args:
103
+ sequence: Validated amino acid sequence string.
104
+ warmup_runs: Initial untimed runs to warm CPU caches and instruction paths.
105
+ benchmark_runs: Timed repeated runs to average out timing jitter.
106
+
107
+ Returns:
108
+ BaselineProfileResult with detailed latency, throughput, and memory stats.
109
+ """
110
+ clean_seq = sequence.strip().upper()
111
+ seq_len = len(clean_seq)
112
+
113
+ if not self._is_loaded:
114
+ self.load_model()
115
+
116
+ if self._using_mock:
117
+ return self._benchmark_mock(clean_seq, seq_len)
118
+
119
+ # Warmup CPU caches
120
+ inputs = self.tokenizer(clean_seq, return_tensors="pt")
121
+ inputs = {k: v.to(self.device) for k, v in inputs.items()}
122
+
123
+ with torch.no_grad():
124
+ for _ in range(warmup_runs):
125
+ _ = self.model(**inputs)
126
+
127
+ # Memory profiling with tracemalloc
128
+ tracemalloc.start()
129
+ start_mem_peak = tracemalloc.get_traced_memory()[1]
130
+
131
+ # Tokenization timing
132
+ t_tok_start = time.perf_counter()
133
+ tokenized_inputs = self.tokenizer(clean_seq, return_tensors="pt")
134
+ tokenized_inputs = {k: v.to(self.device) for k, v in tokenized_inputs.items()}
135
+ t_tok_end = time.perf_counter()
136
+ tokenization_ms = (t_tok_end - t_tok_start) * 1000.0
137
+
138
+ # Timed forward passes
139
+ forward_times = []
140
+ last_hidden_state = None
141
+
142
+ with torch.no_grad():
143
+ for _ in range(max(1, benchmark_runs)):
144
+ t_fwd_start = time.perf_counter()
145
+ outputs = self.model(**tokenized_inputs)
146
+ t_fwd_end = time.perf_counter()
147
+ forward_times.append((t_fwd_end - t_fwd_start) * 1000.0)
148
+ last_hidden_state = outputs.last_hidden_state
149
+
150
+ current_mem, peak_mem = tracemalloc.get_traced_memory()
151
+ tracemalloc.stop()
152
+
153
+ peak_memory_mb = max(0.01, (peak_mem - start_mem_peak) / (1024.0 * 1024.0))
154
+
155
+ # Average forward pass duration
156
+ avg_forward_ms = sum(forward_times) / len(forward_times)
157
+ total_latency_ms = tokenization_ms + avg_forward_ms
158
+
159
+ # Throughput calculations
160
+ ms_per_residue = total_latency_ms / max(1, seq_len)
161
+ tokens_per_sec = (seq_len / (total_latency_ms / 1000.0)) if total_latency_ms > 0 else 0.0
162
+
163
+ # Tensor memory footprint
164
+ tensor_bytes = 0
165
+ tensor_shape = (1, seq_len, 320)
166
+ if last_hidden_state is not None:
167
+ tensor_bytes = last_hidden_state.element_size() * last_hidden_state.nelement()
168
+ tensor_shape = tuple(last_hidden_state.shape)
169
+
170
+ return BaselineProfileResult(
171
+ sequence_length=seq_len,
172
+ tokenization_ms=round(tokenization_ms, 3),
173
+ forward_pass_ms=round(avg_forward_ms, 3),
174
+ total_cpu_latency_ms=round(total_latency_ms, 3),
175
+ ms_per_residue=round(ms_per_residue, 4),
176
+ tokens_per_sec=round(tokens_per_sec, 2),
177
+ peak_memory_mb=round(peak_memory_mb, 2),
178
+ tensor_bytes=tensor_bytes,
179
+ tensor_shape=tensor_shape,
180
+ model_name=self.model_name,
181
+ is_mock=False,
182
+ metadata={"benchmark_runs": benchmark_runs, "device": str(self.device)},
183
+ )
184
+
185
+ def _benchmark_mock(self, clean_seq: str, seq_len: int) -> BaselineProfileResult:
186
+ """
187
+ Synthesizes realistic CPU unaccelerated baseline latency metrics.
188
+ On an average modern x86 CPU, 8M ESM-2 runs at roughly 0.45 - 1.2 ms per residue
189
+ with quadratic attention growth overhead at longer sequences.
190
+ """
191
+ # Baseline CPU performance characteristics
192
+ tok_ms = 0.5 + (seq_len * 0.005)
193
+ # O(N) feedforward + O(N^2) unaccelerated attention component
194
+ fwd_ms = (seq_len * 0.55) + (0.0004 * (seq_len**2))
195
+ total_ms = tok_ms + fwd_ms
196
+
197
+ ms_per_res = total_ms / seq_len
198
+ tokens_per_sec = seq_len / (total_ms / 1000.0)
199
+
200
+ # Emulate CPU processing delay (proportional but bounded for quick benchmark execution)
201
+ emulated_delay = min(0.08, total_ms / 2000.0)
202
+ time.sleep(emulated_delay)
203
+
204
+ hidden_dim = 320 # ESM2-t6-8M hidden dimension
205
+ tensor_shape = (1, seq_len + 2, hidden_dim) # includes CLS and EOS tokens
206
+ tensor_bytes = (seq_len + 2) * hidden_dim * 4 # float32 = 4 bytes
207
+ peak_mem_mb = 12.0 + (seq_len * 0.02)
208
+
209
+ return BaselineProfileResult(
210
+ sequence_length=seq_len,
211
+ tokenization_ms=round(tok_ms, 3),
212
+ forward_pass_ms=round(fwd_ms, 3),
213
+ total_cpu_latency_ms=round(total_ms, 3),
214
+ ms_per_residue=round(ms_per_res, 4),
215
+ tokens_per_sec=round(tokens_per_sec, 2),
216
+ peak_memory_mb=round(peak_mem_mb, 2),
217
+ tensor_bytes=tensor_bytes,
218
+ tensor_shape=tensor_shape,
219
+ model_name=f"{self.model_name} (CPU Simulated)",
220
+ is_mock=True,
221
+ metadata={"device": "cpu-simulated"},
222
+ )
bus/message_bus.py ADDED
@@ -0,0 +1,101 @@
1
+ """
2
+ Agentic BioNeMo - Asynchronous Swarm Message Bus
3
+ Manages inter-agent communication, topic routing, debate history, and real-time streaming hooks.
4
+ """
5
+ from typing import Callable, List, Dict, Any, Optional
6
+ import time
7
+ import logging
8
+ from datetime import datetime
9
+ from src.models import CouncilMessage
10
+
11
+ logger = logging.getLogger("SwarmMessageBus")
12
+
13
+ class SwarmMessageBus:
14
+ def __init__(self):
15
+ self.subscribers: Dict[str, List[Callable[[CouncilMessage], None]]] = {}
16
+ self.global_listeners: List[Callable[[CouncilMessage], None]] = []
17
+ self.history: List[CouncilMessage] = []
18
+ self.active_vetoes: List[CouncilMessage] = []
19
+ self.clearances: List[CouncilMessage] = []
20
+
21
+ def subscribe(self, topic: str, callback: Callable[[CouncilMessage], None]):
22
+ """Subscribe to a specific message intent/topic."""
23
+ if topic not in self.subscribers:
24
+ self.subscribers[topic] = []
25
+ self.subscribers[topic].append(callback)
26
+
27
+ def add_global_listener(self, callback: Callable[[CouncilMessage], None]):
28
+ """Subscribe to all messages published across the entire swarm."""
29
+ self.global_listeners.append(callback)
30
+
31
+ def publish(
32
+ self,
33
+ agent_id: str,
34
+ persona_name: str,
35
+ avatar: str,
36
+ intent: str,
37
+ content: str,
38
+ metadata: Optional[Dict[str, Any]] = None
39
+ ) -> CouncilMessage:
40
+ """Publish a structured council message to the swarm."""
41
+ now_str = datetime.now().strftime("%H:%M:%S")
42
+ msg = CouncilMessage(
43
+ agent_id=agent_id,
44
+ persona_name=persona_name,
45
+ avatar=avatar,
46
+ intent=intent,
47
+ content=content,
48
+ metadata=metadata or {},
49
+ timestamp=time.time(),
50
+ timestamp_str=now_str
51
+ )
52
+
53
+ self.history.append(msg)
54
+
55
+ if intent == "VETO":
56
+ self.active_vetoes.append(msg)
57
+ elif intent == "CLEARANCE":
58
+ self.clearances.append(msg)
59
+
60
+ # Notify topic-specific subscribers
61
+ if intent in self.subscribers:
62
+ for cb in self.subscribers[intent]:
63
+ try:
64
+ cb(msg)
65
+ except Exception as e:
66
+ logger.exception(f"Error in subscriber callback for {intent}: {e}")
67
+
68
+ # Notify global listeners (e.g. UI WebSocket / SSE stream)
69
+ for g_cb in self.global_listeners:
70
+ try:
71
+ g_cb(msg)
72
+ except Exception as e:
73
+ logger.exception(f"Error in global listener callback: {e}")
74
+
75
+ return msg
76
+
77
+ def get_council_summary(self) -> Dict[str, Any]:
78
+ """Returns consolidated council statistics."""
79
+ return {
80
+ "total_messages": len(self.history),
81
+ "veto_count": len(self.active_vetoes),
82
+ "clearance_count": len(self.clearances),
83
+ "last_action": self.history[-1].intent if self.history else "IDLE",
84
+ "recent_dialogues": [
85
+ {
86
+ "agent_id": m.agent_id,
87
+ "persona": m.persona_name,
88
+ "avatar": m.avatar,
89
+ "intent": m.intent,
90
+ "content": m.content,
91
+ "time": m.timestamp_str
92
+ }
93
+ for m in self.history[-12:]
94
+ ]
95
+ }
96
+
97
+ def clear(self):
98
+ """Reset the message bus for a new campaign."""
99
+ self.history.clear()
100
+ self.active_vetoes.clear()
101
+ self.clearances.clear()