diff --git a/audit/__init__.py b/audit/__init__.py index 34a8bfdd78..ee3faf1a69 100644 --- a/audit/__init__.py +++ b/audit/__init__.py @@ -3,6 +3,7 @@ This module provides a thin adapter around `drug_discovery.compliance.audit_trail` so teams can import `audit` as a separate module. """ + from .audit_adapter import ComplianceAuditAdapter __all__ = ["ComplianceAuditAdapter"] diff --git a/audit/audit_adapter.py b/audit/audit_adapter.py index 278fde255d..9898e8c33c 100644 --- a/audit/audit_adapter.py +++ b/audit/audit_adapter.py @@ -1,13 +1,13 @@ """Audit adapter exposing the project's audit trail from a top-level package.""" + from __future__ import annotations -from typing import Any, Dict +from typing import Any from drug_discovery.compliance.audit_trail import ( - ComplianceAuditLogger, AuditTrail, - AuditEventType, ComplianceAuditEntry, + ComplianceAuditLogger, ) @@ -20,11 +20,15 @@ def __init__(self, trail: AuditTrail | None = None): def log_screen(self, smiles: str, compound_id: str | None = None, user_id: str = "system") -> ComplianceAuditEntry: return self.logger.log_compound_screened(smiles=smiles, compound_id=compound_id, user_id=user_id) - def log_prediction(self, compound_id: str, smiles: str, predictions: Dict[str, float], user_id: str = "system") -> ComplianceAuditEntry: - return self.logger.log_toxicity_prediction(compound_id=compound_id, smiles=smiles, predictions=predictions, user_id=user_id) + def log_prediction( + self, compound_id: str, smiles: str, predictions: dict[str, float], user_id: str = "system" + ) -> ComplianceAuditEntry: + return self.logger.log_toxicity_prediction( + compound_id=compound_id, smiles=smiles, predictions=predictions, user_id=user_id + ) def verify(self) -> bool: return self.logger.verify_integrity() - def export(self) -> Dict[str, Any]: + def export(self) -> dict[str, Any]: return self.logger.export_report() diff --git a/backend/api/client_gateway.py b/backend/api/client_gateway.py index 24645f24f9..2369c94c54 100644 --- a/backend/api/client_gateway.py +++ b/backend/api/client_gateway.py @@ -1,11 +1,11 @@ -from fastapi import FastAPI, UploadFile, File, Form -from pydantic import BaseModel, EmailStr -from typing import List, Optional import os -import uuid -import shutil -import asyncio import random +import shutil +import uuid + +from fastapi import FastAPI, File, Form, UploadFile +from pydantic import BaseModel + from drug_discovery.commercial.fda_drug_matcher import CommercialDrugMapper from zane_apex_entrypoint import execute_zane_pipeline @@ -15,6 +15,7 @@ mapper = CommercialDrugMapper() mapper.load_fda_orange_book("data/fda_orange_book.csv") + class Compound(BaseModel): smiles: str dosage: str @@ -22,17 +23,20 @@ class Compound(BaseModel): purpose: str toxicity_level: str + class CommercialMatch(BaseModel): closest_drug: str similarity: float commercial_dose: str - extra_compounds: List[str] - missing_compounds: List[str] + extra_compounds: list[str] + missing_compounds: list[str] + class TherapeuticBlueprint(BaseModel): - compounds: List[Compound] + compounds: list[Compound] commercial_match: CommercialMatch + @app.post("/api/v1/generate_blueprint", response_model=TherapeuticBlueprint) async def generate_blueprint( name: str = Form(...), @@ -44,7 +48,7 @@ async def generate_blueprint( current_treatments: str = Form(""), lifestyle: str = Form(...), hereditary_problems: str = Form(""), - health_report: UploadFile = File(...) + health_report: UploadFile = File(...), ): """ Triggers the ZANE Zero-Mortality engine and returns a detailed therapeutic blueprint. @@ -54,7 +58,7 @@ async def generate_blueprint( temp_id = str(uuid.uuid4()) upload_dir = f"temp_uploads/{temp_id}" os.makedirs(upload_dir, exist_ok=True) - + report_path = os.path.join(upload_dir, health_report.filename) with open(report_path, "wb") as buffer: shutil.copyfileobj(health_report.file, buffer) @@ -68,50 +72,51 @@ async def generate_blueprint( "location": location, "treatments": current_treatments, "lifestyle": lifestyle, - "hereditary": hereditary_problems + "hereditary": hereditary_problems, } await execute_zane_pipeline(report_path, target_purpose, metadata=metadata) - + # 3. Generate Multi-Compound Blueprint (10-20 compounds) num_compounds = random.randint(10, 20) mock_compounds = [] - + # Primary compound primary_smiles = "CC1=C(C=C(C=C1)NC(=O)C2=CC=C(C=C2)CN3CCN(CC3)C)NC4=NC=CC(=N4)C5=CN=CC=C5" - mock_compounds.append(Compound( - smiles=primary_smiles, - dosage="14.5mg", - timing="08:30 AM", - purpose=f"Primary inhibitor for {target_purpose}", - toxicity_level="Ultra-Low (0.02 LD50)" - )) - + mock_compounds.append( + Compound( + smiles=primary_smiles, + dosage="14.5mg", + timing="08:30 AM", + purpose=f"Primary inhibitor for {target_purpose}", + toxicity_level="Ultra-Low (0.02 LD50)", + ) + ) + # Adjuvant compounds for i in range(num_compounds - 1): - mock_compounds.append(Compound( - smiles=f"SMILES_ADJ_{i}_{uuid.uuid4().hex[:6]}", - dosage=f"{random.uniform(1, 10):.1f}mg", - timing=f"{random.randint(8, 22):02d}:00", - purpose="Metabolic synergy / Adjuvant", - toxicity_level="Non-toxic" - )) - + mock_compounds.append( + Compound( + smiles=f"SMILES_ADJ_{i}_{uuid.uuid4().hex[:6]}", + dosage=f"{random.uniform(1, 10):.1f}mg", + timing=f"{random.randint(8, 22):02d}:00", + purpose="Metabolic synergy / Adjuvant", + toxicity_level="Non-toxic", + ) + ) + # 4. Find Commercial Match for the primary compound comm_match_data = mapper.find_closest_commercial_match(primary_smiles) - + # 5. Compare with multi-compound ZANE drug comp_analysis = mapper.compare_compounds([c.dict() for c in mock_compounds], comm_match_data) - + comm_match = CommercialMatch( - closest_drug=comm_match_data['closest_drug'], - similarity=comm_match_data['similarity'], - commercial_dose=comm_match_data['commercial_dose'], - extra_compounds=comp_analysis['extra_compounds'], - missing_compounds=comp_analysis['missing_compounds'] + closest_drug=comm_match_data["closest_drug"], + similarity=comm_match_data["similarity"], + commercial_dose=comm_match_data["commercial_dose"], + extra_compounds=comp_analysis["extra_compounds"], + missing_compounds=comp_analysis["missing_compounds"], ) - + # 6. Final Blueprint - return TherapeuticBlueprint( - compounds=mock_compounds, - commercial_match=comm_match - ) + return TherapeuticBlueprint(compounds=mock_compounds, commercial_match=comm_match) diff --git a/clinical/chronobiology/circadian_dosing.py b/clinical/chronobiology/circadian_dosing.py index 56b35322e7..aa4e5a5285 100644 --- a/clinical/chronobiology/circadian_dosing.py +++ b/clinical/chronobiology/circadian_dosing.py @@ -1,19 +1,21 @@ +import logging + import numpy as np import pandas as pd from scipy.optimize import least_squares -from typing import Dict, Any, Optional -import logging logger = logging.getLogger(__name__) + class CircadianDosingOptimizer: """ - Optimizes drug dosing schedules based on patient-specific circadian rhythms + Optimizes drug dosing schedules based on patient-specific circadian rhythms derived from wearable telemetry. """ + def __init__(self): - self.telemetry_data: Optional[pd.DataFrame] = None - self.circadian_params: Dict[str, float] = {} + self.telemetry_data: pd.DataFrame | None = None + self.circadian_params: dict[str, float] = {} def ingest_wearable_telemetry(self, timeseries_csv: str) -> None: """ @@ -22,62 +24,64 @@ def ingest_wearable_telemetry(self, timeseries_csv: str) -> None: """ try: df = pd.read_csv(timeseries_csv) - df['timestamp'] = pd.to_datetime(df['timestamp']) - df['hour'] = df['timestamp'].dt.hour + df['timestamp'].dt.minute / 60.0 + df["timestamp"] = pd.to_datetime(df["timestamp"]) + df["hour"] = df["timestamp"].dt.hour + df["timestamp"].dt.minute / 60.0 self.telemetry_data = df - + # Fit a cosinor model to body temperature to find the acrophase (peak) # T(t) = M + A * cos(2*pi*t/24 - phi) # Dim Light Melatonin Onset (DLMO) is typically ~7 hours before the temperature nadir self._fit_circadian_model(df) - + except Exception as e: - logger.error(f"Error ingesting telemetry: {str(e)}") + logger.error(f"Error ingesting telemetry: {e!s}") # Default fallback for healthy adult self.circadian_params = {"acrophase": 16.0, "mesor": 37.0, "amplitude": 0.5} def _fit_circadian_model(self, df: pd.DataFrame): """Fits a cosinor model to the telemetry data.""" - t = df['hour'].values - y = df['body_temperature'].values if 'body_temperature' in df else df['heart_rate'].values - + t = df["hour"].values + y = df["body_temperature"].values if "body_temperature" in df else df["heart_rate"].values + def model(params, t): mesor, amp, phi = params return mesor + amp * np.cos(2 * np.pi * t / 24 - phi) - + def residuals(params, t, y): return model(params, t) - y - + # Initial guess x0 = [np.mean(y), np.std(y), 0] res = least_squares(residuals, x0, args=(t, y)) - + self.circadian_params = { "mesor": res.x[0], "amplitude": res.x[1], - "acrophase": res.x[2] % (2 * np.pi) * 24 / (2 * np.pi) + "acrophase": res.x[2] % (2 * np.pi) * 24 / (2 * np.pi), } logger.info(f"Circadian phase detected: Acrophase at {self.circadian_params['acrophase']:.2f}h") def calculate_optimal_tmax(self, target_receptor_peak_hour: float = 8.0) -> float: """ Calculates the optimal hour to administer the drug. - peak_plasma_concentration (Cmax) should align with the peak circadian expression + peak_plasma_concentration (Cmax) should align with the peak circadian expression of the target disease receptor. """ # Adjust target peak based on patient's specific phase shift # Standard healthy acrophase is approx 16:00 (4 PM) phase_shift = self.circadian_params.get("acrophase", 16.0) - 16.0 - + # Shift the target receptor peak by the patient's individual clock shift individualized_target_hour = (target_receptor_peak_hour + phase_shift) % 24 - + # If the drug takes 'absorption_delay' hours to reach Tmax - absorption_delay = 2.0 # Assume 2 hours for standard oral delivery - + absorption_delay = 2.0 # Assume 2 hours for standard oral delivery + optimal_dosing_time = (individualized_target_hour - absorption_delay) % 24 - - logger.info(f"Optimal dosing time calculated: {optimal_dosing_time:.2f}h " - f"to reach peak at {individualized_target_hour:.2f}h") - + + logger.info( + f"Optimal dosing time calculated: {optimal_dosing_time:.2f}h " + f"to reach peak at {individualized_target_hour:.2f}h" + ) + return optimal_dosing_time diff --git a/clinical/digital_twin/epigenetic_profiler.py b/clinical/digital_twin/epigenetic_profiler.py index f44ab72f59..b671a02efa 100644 --- a/clinical/digital_twin/epigenetic_profiler.py +++ b/clinical/digital_twin/epigenetic_profiler.py @@ -1,16 +1,17 @@ +import logging + import numpy as np -from sklearn.linear_model import LinearRegression import pandas as pd -from typing import Optional -import logging logger = logging.getLogger(__name__) + class EpigeneticAgeCalculator: """ - Calculates biological age using DNA methylation data and adjusts therapeutic + Calculates biological age using DNA methylation data and adjusts therapeutic dosing to prevent overdose in senescent biological systems. """ + def __init__(self): # Pre-trained Horvath Clock coefficients (simplified mock subset) # Real Horvath clock uses 353 specific CpG sites @@ -19,7 +20,7 @@ def __init__(self): "cg00374713": 0.12, "cg00864867": -0.08, "cg01234567": 0.22, - "intercept": 0.65 + "intercept": 0.65, } def calculate_horvath_clock(self, methylation_array_path: str) -> float: @@ -29,48 +30,52 @@ def calculate_horvath_clock(self, methylation_array_path: str) -> float: try: # Expecting a CSV with 'cpg_id' and 'beta_value' (0 to 1) df = pd.read_csv(methylation_array_path) - + # Linear combination of methylation levels log_age = self.horvath_coefficients["intercept"] for cpg, coeff in self.horvath_coefficients.items(): - if cpg == "intercept": continue - - beta_value = df.loc[df['cpg_id'] == cpg, 'beta_value'] + if cpg == "intercept": + continue + + beta_value = df.loc[df["cpg_id"] == cpg, "beta_value"] if not beta_value.empty: log_age += coeff * beta_value.values[0] - + # Inverse of the age transformation used by Horvath # (Simplified: assuming biological_age = exp(log_age)) biological_age = np.exp(log_age) - + logger.info(f"Epigenetic analysis complete. Biological age: {biological_age:.1f} years.") return float(biological_age) - + except Exception as e: - logger.error(f"Failed to calculate epigenetic age: {str(e)}") - return 45.0 # Default adult age + logger.error(f"Failed to calculate epigenetic age: {e!s}") + return 45.0 # Default adult age - def adjust_dosing_for_senescence(self, base_dose: float, biological_age: float, chronological_age: float = 50.0) -> float: + def adjust_dosing_for_senescence( + self, base_dose: float, biological_age: float, chronological_age: float = 50.0 + ) -> float: """ - Dynamically scales the therapeutic dose if the patient's epigenetic age + Dynamically scales the therapeutic dose if the patient's epigenetic age indicates severe cellular senescence. """ # Calculate the aging acceleration factor age_acceleration = biological_age - chronological_age - + # If the patient is "biologically older" than their years, reduce dose # due to expected decline in metabolic capacity and cellular repair. reduction_factor = 1.0 - + if age_acceleration > 5.0: # Reduce dose by 1% for every year of acceleration beyond 5 years excess_aging = age_acceleration - 5.0 reduction_factor = max(0.5, 1.0 - (excess_aging * 0.02)) - + adjusted_dose = base_dose * reduction_factor - + if reduction_factor < 1.0: - logger.warning(f"Dose adjusted for senescence. Factor: {reduction_factor:.2f}. " - f"New dose: {adjusted_dose:.2f}") - + logger.warning( + f"Dose adjusted for senescence. Factor: {reduction_factor:.2f}. " f"New dose: {adjusted_dose:.2f}" + ) + return float(adjusted_dose) diff --git a/clinical/digital_twin/lab_report_parser.py b/clinical/digital_twin/lab_report_parser.py index b251fcac66..65e40c017f 100644 --- a/clinical/digital_twin/lab_report_parser.py +++ b/clinical/digital_twin/lab_report_parser.py @@ -1,34 +1,38 @@ -import pandas as pd -from pydantic import BaseModel, Field -from typing import Optional, Dict, Any, List -import pdfplumber -import re import logging +import re +from typing import Any + +import pdfplumber +from pydantic import BaseModel, Field # Configure logging logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) + class PatientHealthState(BaseModel): """ Represents the mathematical constraints of a patient's health state. """ + patient_id: str egfr: float = Field(..., description="Estimated Glomerular Filtration Rate (mL/min/1.73m2)") ast: float = Field(..., description="Aspartate Aminotransferase (U/L)") alt: float = Field(..., description="Alanine Aminotransferase (U/L)") blood_ph: float = Field(default=7.4, description="Systemic blood pH") sodium_level: float = Field(..., description="Serum sodium level (mmol/L)") - allergies: List[str] = Field(default_factory=list) - conditions: List[str] = Field(default_factory=list) - organ_viability: Dict[str, str] = Field(default_factory=dict) + allergies: list[str] = Field(default_factory=list) + conditions: list[str] = Field(default_factory=list) + organ_viability: dict[str, str] = Field(default_factory=dict) + class LabReportIngestor: """ Ingests and parses unstructured health reports to extract biological constraints. """ + def __init__(self): - self.constraints: Dict[str, Any] = {} + self.constraints: dict[str, Any] = {} def parse_metabolic_panel(self, pdf_path: str) -> PatientHealthState: """ @@ -40,8 +44,8 @@ def parse_metabolic_panel(self, pdf_path: str) -> PatientHealthState: for page in pdf.pages: text += page.extract_text() or "" except Exception as e: - logger.error(f"Failed to parse PDF {pdf_path}: {str(e)}") - raise ValueError(f"Could not read PDF: {str(e)}") + logger.error(f"Failed to parse PDF {pdf_path}: {e!s}") + raise ValueError(f"Could not read PDF: {e!s}") # Extract biomarkers using regex egfr = self._extract_value(text, r"eGFR[:\s]+(\d+\.?\d*)") @@ -54,17 +58,21 @@ def parse_metabolic_panel(self, pdf_path: str) -> PatientHealthState: logger.warning("Some critical biomarkers were not found in the report. Using defaults where appropriate.") return PatientHealthState( - patient_id="PID-" + re.search(r"Patient ID[:\s]+(\w+)", text, re.I).group(1) if re.search(r"Patient ID[:\s]+(\w+)", text, re.I) else "UNKNOWN", + patient_id=( + "PID-" + re.search(r"Patient ID[:\s]+(\w+)", text, re.I).group(1) + if re.search(r"Patient ID[:\s]+(\w+)", text, re.I) + else "UNKNOWN" + ), egfr=egfr if egfr is not None else 90.0, ast=ast if ast is not None else 25.0, alt=alt if alt is not None else 25.0, blood_ph=ph, sodium_level=sodium if sodium is not None else 140.0, allergies=re.findall(r"Allergy[:\s]+(\w+)", text, re.I), - conditions=re.findall(r"Condition[:\s]+(\w+)", text, re.I) + conditions=re.findall(r"Condition[:\s]+(\w+)", text, re.I), ) - def _extract_value(self, text: str, pattern: str) -> Optional[float]: + def _extract_value(self, text: str, pattern: str) -> float | None: match = re.search(pattern, text, re.IGNORECASE) if match: try: @@ -73,7 +81,7 @@ def _extract_value(self, text: str, pattern: str) -> Optional[float]: return None return None - def evaluate_organ_viability(self, state: PatientHealthState) -> Dict[str, Any]: + def evaluate_organ_viability(self, state: PatientHealthState) -> dict[str, Any]: """ Flags severe organ impairment and locks parameters into a constraint dictionary. """ diff --git a/clinical/digital_twin/microbiome_metabolomics.py b/clinical/digital_twin/microbiome_metabolomics.py index 729c287dfb..1a67cf55c0 100644 --- a/clinical/digital_twin/microbiome_metabolomics.py +++ b/clinical/digital_twin/microbiome_metabolomics.py @@ -1,21 +1,22 @@ -import pandas as pd -import networkx as nx -from Bio import SeqIO -from typing import Dict, List, Optional import logging + +from Bio import SeqIO from rdkit import Chem -from rdkit.Chem import AllChem logger = logging.getLogger(__name__) + class MicrobiomeToxicityVeto(Exception): """Exception raised when a drug is predicted to be metabolized into a toxic byproduct by gut flora.""" + pass + class PharmacobiomiomicEngine: """ Analyzes the interaction between a patient's gut microbiome and drug candidates. """ + def __init__(self): # Known microbial enzymatic reactions (simplified mapping) # In a real system, this would be a large database of metabolic pathways @@ -23,18 +24,18 @@ def __init__(self): "azoreductase": ["Bacteroides", "Clostridium", "Enterococcus"], "beta-glucuronidase": ["Escherichia coli", "Bacteroides vulgatus"], "nitroreductase": ["Bacteroides fragilis"], - "sulfatase": ["Peptostreptococcus"] + "sulfatase": ["Peptostreptococcus"], } - + # Chemical fragments targeted by these enzymes self.reactive_fragments = { "azoreductase": "N=N", - "beta-glucuronidase": "OC1OC(C(O)C(O)C1O)C(=O)O", # Glucuronide fragment + "beta-glucuronidase": "OC1OC(C(O)C(O)C1O)C(=O)O", # Glucuronide fragment "nitroreductase": "[N+](=O)[O-]", - "sulfatase": "OS(=O)(=O)O" + "sulfatase": "OS(=O)(=O)O", } - def parse_metagenomic_seq(self, fastq_path: str) -> Dict[str, float]: + def parse_metagenomic_seq(self, fastq_path: str) -> dict[str, float]: """ Parses metagenomic sequencing data to identify abundance of bacterial strains. """ @@ -47,24 +48,24 @@ def parse_metagenomic_seq(self, fastq_path: str) -> Dict[str, float]: record_count += 1 # Mock taxonomic classification logic seq_str = str(record.seq) - if "GATC" in seq_str: # Mock marker for Bacteroides + if "GATC" in seq_str: # Mock marker for Bacteroides abundance_profile["Bacteroides"] = abundance_profile.get("Bacteroides", 0) + 1 - if "ATGC" in seq_str: # Mock marker for E. coli + if "ATGC" in seq_str: # Mock marker for E. coli abundance_profile["Escherichia coli"] = abundance_profile.get("Escherichia coli", 0) + 1 - + # Normalize to relative abundance if record_count > 0: for genus in abundance_profile: abundance_profile[genus] /= record_count - + except Exception as e: - logger.error(f"Failed to parse metagenomic data {fastq_path}: {str(e)}") + logger.error(f"Failed to parse metagenomic data {fastq_path}: {e!s}") # Return a default profile if parsing fails return {"Bacteroides": 0.4, "Escherichia coli": 0.1} - + return abundance_profile - def predict_microbial_cleavage(self, smiles: str, microbiome_profile: Dict[str, float]) -> bool: + def predict_microbial_cleavage(self, smiles: str, microbiome_profile: dict[str, float]) -> bool: """ Maps the drug's SMILES against known bacterial enzymatic actions. Raises MicrobiomeToxicityVeto if high-risk metabolism is predicted. @@ -79,12 +80,14 @@ def predict_microbial_cleavage(self, smiles: str, microbiome_profile: Dict[str, # Check if the patient has the bacteria that produce this enzyme relevant_strains = self.microbial_enzymes.get(enzyme, []) abundance = sum(microbiome_profile.get(strain, 0) for strain in relevant_strains) - - if abundance > 0.15: # Threshold for clinical significance - logger.warning(f"VETO: Drug contains {enzyme} substrate and patient has high abundance of {relevant_strains}") + + if abundance > 0.15: # Threshold for clinical significance + logger.warning( + f"VETO: Drug contains {enzyme} substrate and patient has high abundance of {relevant_strains}" + ) raise MicrobiomeToxicityVeto( f"Molecule contains fragment {pattern} which is likely to be cleaved by {enzyme} " f"present in patient's microbiome (abundance: {abundance:.2%})." ) - + return True diff --git a/clinical/safety/dynamic_polypharmacy_ddi.py b/clinical/safety/dynamic_polypharmacy_ddi.py index e672c78988..a2b6f61830 100644 --- a/clinical/safety/dynamic_polypharmacy_ddi.py +++ b/clinical/safety/dynamic_polypharmacy_ddi.py @@ -1,29 +1,27 @@ -import networkx as nx -import numpy as np -from scipy.integrate import odeint -from typing import List, Dict, Any, Optional import logging + from rdkit import Chem -from rdkit.Chem import Descriptors logger = logging.getLogger(__name__) + class DynamicDDINetwork: """ - Simulates competitive inhibition and metabolic pathway interference between + Simulates competitive inhibition and metabolic pathway interference between AI-generated drugs and existing patient prescriptions. """ + def __init__(self): self.active_meds = [] self.cyp_map = { "Warfarin": ["CYP2C9", "CYP1A2"], "Atorvastatin": ["CYP3A4"], "Clopidogrel": ["CYP2C19"], - "Metoprolol": ["CYP2D6"] + "Metoprolol": ["CYP2D6"], } - self.enzyme_kinetics = {} # Km, Vmax for different CYPs + self.enzyme_kinetics = {} # Km, Vmax for different CYPs - def load_active_prescriptions(self, ehr_medication_list: List[str]): + def load_active_prescriptions(self, ehr_medication_list: list[str]): """ Ingests patient's current drugs and maps metabolic pathways. """ @@ -32,29 +30,35 @@ def load_active_prescriptions(self, ehr_medication_list: List[str]): def simulate_competitive_inhibition(self, generated_smiles: str) -> bool: """ - Mathematically simulates enzyme competition to prevent fatal overdose + Mathematically simulates enzyme competition to prevent fatal overdose buildup of current prescriptions. """ mol = Chem.MolFromSmiles(generated_smiles) - if not mol: return False + if not mol: + return False # Mock prediction of generated drug's primary metabolic pathway # (In practice, this would use a deep learning CYP selectivity model) - predicted_pathway = "CYP3A4" # Most common - + predicted_pathway = "CYP3A4" # Most common + for med in self.active_meds: med_pathways = self.cyp_map.get(med, []) if predicted_pathway in med_pathways: # Fatal Interaction Potential Detected (e.g. inhibiting CYP2C9 while on Warfarin) if predicted_pathway == "CYP2C9" and "Warfarin" in self.active_meds: - logger.error(f"FATAL DDI: Generated drug inhibits {predicted_pathway} while patient is on Warfarin.") + logger.error( + f"FATAL DDI: Generated drug inhibits {predicted_pathway} while patient is on Warfarin." + ) return True - + # Check for CYP3A4 bottleneck - if predicted_pathway == "CYP3A4" and len([m for m in self.active_meds if "CYP3A4" in self.cyp_map.get(m, [])]) > 2: - logger.warning(f"Lethal Polypharmacy Risk: CYP3A4 metabolic bottleneck detected.") + if ( + predicted_pathway == "CYP3A4" + and len([m for m in self.active_meds if "CYP3A4" in self.cyp_map.get(m, [])]) > 2 + ): + logger.warning("Lethal Polypharmacy Risk: CYP3A4 metabolic bottleneck detected.") return True - + return False def pK_ode_system(self, y, t, drug_inflow, enzyme_count): @@ -65,9 +69,9 @@ def pK_ode_system(self, y, t, drug_inflow, enzyme_count): C_gen, C_ex = y Km_gen, Vmax_gen = 1.0, 10.0 Km_ex, Vmax_ex = 1.5, 8.0 - + # Competitive inhibition terms - dC_gen = drug_inflow - (Vmax_gen * C_gen) / (Km_gen * (1 + C_ex/Km_ex) + C_gen) - dC_ex = - (Vmax_ex * C_ex) / (Km_ex * (1 + C_gen/Km_gen) + C_ex) - + dC_gen = drug_inflow - (Vmax_gen * C_gen) / (Km_gen * (1 + C_ex / Km_ex) + C_gen) + dC_ex = -(Vmax_ex * C_ex) / (Km_ex * (1 + C_gen / Km_gen) + C_ex) + return [dC_gen, dC_ex] diff --git a/clinical/safety/hla_hypersensitivity_veto.py b/clinical/safety/hla_hypersensitivity_veto.py index b5aa220dcf..43aaa47cde 100644 --- a/clinical/safety/hla_hypersensitivity_veto.py +++ b/clinical/safety/hla_hypersensitivity_veto.py @@ -1,49 +1,49 @@ -import torch -import torch.nn as nn -import pandas as pd import json import logging + +import torch +import torch.nn as nn from rdkit import Chem from rdkit.Chem import AllChem -from typing import List, Dict, Optional logger = logging.getLogger(__name__) + class LethalHypersensitivityVeto(Exception): """Exception raised when a drug is predicted to trigger a lethal HLA-mediated autoimmune reaction.""" + pass + class PersonalizedImmunotoxScreener: """ Evaluates potential for lethal idiosyncratic hypersensitivity reactions (e.g., SJS/TEN) by predicting drug binding to patient-specific HLA alleles. """ + def __init__(self): # Mock geometric deep learning surrogate model for HLA-drug binding # In production, this would be a GNN-based binding affinity predictor self.binding_model = nn.Sequential( - nn.Linear(2048 + 512, 1024), # Drug fingerprint + HLA embedding - nn.ReLU(), - nn.Linear(1024, 1), - nn.Sigmoid() + nn.Linear(2048 + 512, 1024), nn.ReLU(), nn.Linear(1024, 1), nn.Sigmoid() # Drug fingerprint + HLA embedding ) - self.hla_embeddings: Dict[str, torch.Tensor] = {} # Mock embeddings for common alleles + self.hla_embeddings: dict[str, torch.Tensor] = {} # Mock embeddings for common alleles - def ingest_patient_hla_typing(self, hla_json_path: str) -> List[str]: + def ingest_patient_hla_typing(self, hla_json_path: str) -> list[str]: """ Parses the patient's exact HLA class I and II genotypes. """ try: - with open(hla_json_path, 'r') as f: + with open(hla_json_path) as f: data = json.load(f) alleles = data.get("hla_typing", []) logger.info(f"Loaded {len(alleles)} patient HLA alleles: {alleles}") return alleles except Exception as e: - logger.error(f"Failed to parse HLA typing at {hla_json_path}: {str(e)}") - return ["HLA-B*57:01"] # High-risk default if unknown + logger.error(f"Failed to parse HLA typing at {hla_json_path}: {e!s}") + return ["HLA-B*57:01"] # High-risk default if unknown - def predict_hla_drug_complex(self, smiles: str, patient_hla_alleles: List[str]) -> float: + def predict_hla_drug_complex(self, smiles: str, patient_hla_alleles: list[str]) -> float: """ Predicts if the generated molecule will bind into the patient's HLA antigen-presentation groove. """ @@ -59,17 +59,17 @@ def predict_hla_drug_complex(self, smiles: str, patient_hla_alleles: List[str]) for allele in patient_hla_alleles: # Get or generate mock HLA embedding hla_emb = self.hla_embeddings.get(allele, torch.randn(512)) - + # Predict binding risk combined = torch.cat([drug_tensor, hla_emb]) risk_score = self.binding_model(combined).item() - - if risk_score > 0.85: # Threshold for triggering a T-cell response + + if risk_score > 0.85: # Threshold for triggering a T-cell response logger.error(f"LETHAL VETO: High risk of HLA-mediated hypersensitivity with allele {allele}") raise LethalHypersensitivityVeto( f"Molecule predicted to bind into {allele} groove, likely triggering lethal autoimmune reaction." ) - + max_risk = max(max_risk, risk_score) - + return max_risk diff --git a/clinical_intelligence/post_market_edge.py b/clinical_intelligence/post_market_edge.py index a9220ca066..f0e50766f5 100644 --- a/clinical_intelligence/post_market_edge.py +++ b/clinical_intelligence/post_market_edge.py @@ -20,7 +20,7 @@ def hidden_markov_model_state_classification(self, sensor_time_series: np.ndarra Hidden Markov Models (HMM): Deploys continuous-time HMMs to classify the patient's baseline physiological state (e.g., sleep, active, resting). - + When tslearn is unavailable, uses heuristic classifier based on signal statistics. State 0: Resting, 1: Active, 2: Sleep """ @@ -36,17 +36,17 @@ def hidden_markov_model_state_classification(self, sensor_time_series: np.ndarra # Compute local statistics for windowed classification states = [] window_size = max(10, len(sensor_time_series) // 20) - + for i in range(0, len(sensor_time_series), max(1, window_size // 2)): - window = sensor_time_series[i:min(i + window_size, len(sensor_time_series))] + window = sensor_time_series[i : min(i + window_size, len(sensor_time_series))] if len(window) == 0: states.append(1) # Default to Active if window empty continue - + # Signal mean and variance indicate state mean_val = np.mean(window) std_val = np.std(window) - + # Classify based on energy (variance) and mean level if std_val < 0.1: # Low variability -> Resting or Sleep state = 0 if mean_val > 0 else 2 # Resting vs Sleep based on mean @@ -54,17 +54,17 @@ def hidden_markov_model_state_classification(self, sensor_time_series: np.ndarra state = 0 else: # High variability -> Active state = 1 - + # Replicate state for each sample in window for _ in range(min(window_size // 2, len(sensor_time_series) - i)): if len(states) < len(sensor_time_series): states.append(state) - + # Pad or trim to match input length - states = states[:len(sensor_time_series)] + states = states[: len(sensor_time_series)] if len(states) < len(sensor_time_series): states.extend([1] * (len(sensor_time_series) - len(states))) - + return np.array(states) def calculate_rmssd(self, r_peaks_ms: np.ndarray) -> float: @@ -118,7 +118,7 @@ def detect_subclinical_stress(self, hrv_data: np.ndarray) -> dict: Uses an LSTM Autoencoder to detect sub-clinical autonomic nervous system stress, flagging potential hepatotoxicity or cardiotoxicity months before the patient physically feels symptoms. - + When the model is unavailable, uses heuristic anomaly scoring based on HRV statistics. """ # Reshape data for LSTM (samples, timesteps, features) @@ -133,19 +133,16 @@ def detect_subclinical_stress(self, hrv_data: np.ndarray) -> dict: else: mean_val = np.mean(hrv_data) std_val = np.std(hrv_data) - + # Normalize to avoid division by zero - if mean_val == 0: - cv = std_val - else: - cv = std_val / np.abs(mean_val) - + cv = std_val if mean_val == 0 else std_val / np.abs(mean_val) + # Reconstruction error based on regularity # Low CV = regular (good reconstruction) -> low error # High CV = irregular (poor reconstruction) -> high error # Scale to [0, 1] range reconstruction_error = min(1.0, cv / 2.0) - + # Add small noise to simulate imperfect reconstruction reconstruction_error += np.random.normal(0, 0.05) reconstruction_error = np.clip(reconstruction_error, 0, 1) diff --git a/compliance/fda_part11_alcoa.py b/compliance/fda_part11_alcoa.py index 34de3fec85..4d53c9d59c 100644 --- a/compliance/fda_part11_alcoa.py +++ b/compliance/fda_part11_alcoa.py @@ -1,18 +1,22 @@ -import json -import hashlib import functools +import hashlib +import json +from collections.abc import Callable from datetime import datetime -from typing import Any, Callable, Dict, Optional -from sqlalchemy import Column, Integer, String, Text, DateTime, create_engine +from typing import Any + +from sqlalchemy import Column, DateTime, Integer, String, Text, create_engine from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker Base = declarative_base() + class AuditEntry(Base): """SQLAlchemy model for ALCOA+ audit trail.""" - __tablename__ = 'audit_ledger' - + + __tablename__ = "audit_ledger" + id = Column(Integer, primary_key=True) timestamp = Column(DateTime, default=datetime.utcnow, nullable=False) user_id = Column(String(100), nullable=False) @@ -22,12 +26,13 @@ class AuditEntry(Base): current_hash = Column(String(64), nullable=False, unique=True) metadata_json = Column(Text) # For ALCOA+ extra context + class ALCOAAuditLedger: """ - Enforces Attributable, Legible, Contemporaneous, Original, and Accurate (ALCOA+) + Enforces Attributable, Legible, Contemporaneous, Original, and Accurate (ALCOA+) data integrity as per 21 CFR Part 11 and EMA Annex 11. """ - + def __init__(self, db_url: str = "sqlite:///compliance_audit.db"): self.engine = create_engine(db_url) Base.metadata.create_all(self.engine) @@ -40,7 +45,9 @@ def _get_last_hash(self) -> str: session.close() return last_entry.current_hash if last_entry else "0" * 64 - def log_immutable_event(self, user_id: str, function_name: str, params: Dict[str, Any], metadata: Optional[Dict] = None): + def log_immutable_event( + self, user_id: str, function_name: str, params: dict[str, Any], metadata: dict | None = None + ): """ Logs an event with cryptographic chaining to prevent silent tampering. Ensures data is Attributable and Contemporaneous. @@ -50,11 +57,11 @@ def log_immutable_event(self, user_id: str, function_name: str, params: Dict[str previous_hash = self._get_last_hash() param_str = json.dumps(params, sort_keys=True) timestamp = datetime.utcnow() - + # Create payload for hashing (Chaining) payload = f"{previous_hash}|{timestamp.isoformat()}|{user_id}|{function_name}|{param_str}" current_hash = hashlib.sha256(payload.encode()).hexdigest() - + entry = AuditEntry( timestamp=timestamp, user_id=user_id, @@ -62,9 +69,9 @@ def log_immutable_event(self, user_id: str, function_name: str, params: Dict[str parameters=param_str, previous_hash=previous_hash, current_hash=current_hash, - metadata_json=json.dumps(metadata) if metadata else None + metadata_json=json.dumps(metadata) if metadata else None, ) - + session.add(entry) session.commit() return current_hash @@ -74,40 +81,40 @@ def log_immutable_event(self, user_id: str, function_name: str, params: Dict[str finally: session.close() + def requires_e_signature(func: Callable): """ Decorator that forces re-authentication before executing critical functions. Compliance: 21 CFR Part 11.10(j) and 11.50. """ + @functools.wraps(func) def wrapper(*args, **kwargs): # In a production environment, this would interface with an Identity Provider (IdP) # or a session manager to verify recent re-auth for 'Electronic Signature' intent. print(f"[E-SIGNATURE] Verification required for: {func.__name__}") - + # Simulated re-authentication check - authenticated = kwargs.get('__authenticated_esign', False) + authenticated = kwargs.get("__authenticated_esign", False) if not authenticated: raise PermissionError( f"Electronic Signature verification failed for '{func.__name__}'. " "Active session re-authentication is required." ) - + # If authenticated, execute and log - user_id = kwargs.get('user_id', 'unknown_user') + user_id = kwargs.get("user_id", "unknown_user") ledger = ALCOAAuditLedger() - + # Filter sensitive kwargs before logging - log_params = {k: v for k, v in kwargs.items() if k != '__authenticated_esign'} - + log_params = {k: v for k, v in kwargs.items() if k != "__authenticated_esign"} + result = func(*args, **kwargs) - + ledger.log_immutable_event( - user_id=user_id, - function_name=func.__name__, - params=log_params, - metadata={"esignature_verified": True} + user_id=user_id, function_name=func.__name__, params=log_params, metadata={"esignature_verified": True} ) - + return result + return wrapper diff --git a/compliance/privacy_hipaa_gdpr.py b/compliance/privacy_hipaa_gdpr.py index b2e5f8f8c9..2dc0192410 100644 --- a/compliance/privacy_hipaa_gdpr.py +++ b/compliance/privacy_hipaa_gdpr.py @@ -1,6 +1,6 @@ -import pandas as pd import numpy as np -from typing import List, Optional +import pandas as pd + try: from presidio_analyzer import AnalyzerEngine from presidio_anonymizer import AnonymizerEngine @@ -10,12 +10,13 @@ AnalyzerEngine = None AnonymizerEngine = None + class PHISanitizer: """ Enforces HIPAA Safe Harbor and GDPR privacy standards for patient data. Uses Presidio NLP for PII/PHI detection and differential privacy for dataset protection. """ - + def __init__(self): if AnalyzerEngine: self.analyzer = AnalyzerEngine() @@ -24,7 +25,7 @@ def __init__(self): self.analyzer = None self.anonymizer = None - def anonymize_cohort_data(self, df: pd.DataFrame, columns_to_scan: Optional[List[str]] = None) -> pd.DataFrame: + def anonymize_cohort_data(self, df: pd.DataFrame, columns_to_scan: list[str] | None = None) -> pd.DataFrame: """ Redacts 18 HIPAA Safe Harbor identifiers. Scans text columns for Names, Dates, SSNs, MRNs, etc. @@ -35,16 +36,16 @@ def anonymize_cohort_data(self, df: pd.DataFrame, columns_to_scan: Optional[List sanitized_df = df.copy() if columns_to_scan is None: - columns_to_scan = sanitized_df.select_dtypes(include=['object']).columns + columns_to_scan = sanitized_df.select_dtypes(include=["object"]).columns for col in columns_to_scan: sanitized_df[col] = sanitized_df[col].apply(lambda x: self._scrub_text(str(x)) if pd.notnull(x) else x) - + return sanitized_df def _scrub_text(self, text: str) -> str: """Internal helper to analyze and anonymize a single string.""" - results = self.analyzer.analyze(text=text, entities=[], language='en') + results = self.analyzer.analyze(text=text, entities=[], language="en") anonymized_result = self.anonymizer.anonymize( text=text, analyzer_results=results, @@ -56,11 +57,13 @@ def _scrub_text(self, text: str) -> str: "EMAIL_ADDRESS": OperatorConfig("replace", {"new_value": ""}), "US_SSN": OperatorConfig("replace", {"new_value": ""}), "US_PASSPORT": OperatorConfig("replace", {"new_value": ""}), - } + }, ) return anonymized_result.text - def inject_differential_privacy(self, df: pd.DataFrame, epsilon: float = 0.1, columns: Optional[List[str]] = None) -> pd.DataFrame: + def inject_differential_privacy( + self, df: pd.DataFrame, epsilon: float = 0.1, columns: list[str] | None = None + ) -> pd.DataFrame: """ Applies epsilon-differential privacy by adding Laplacian noise to numerical columns. Compliance: GDPR Requirement for non-reversibility. @@ -74,20 +77,22 @@ def inject_differential_privacy(self, df: pd.DataFrame, epsilon: float = 0.1, co sensitivity = dp_df[col].max() - dp_df[col].min() if sensitivity == 0: sensitivity = 1.0 - + beta = sensitivity / epsilon noise = np.random.laplace(0, beta, len(dp_df)) dp_df[col] = dp_df[col] + noise - + return dp_df def sanitize_for_trial_simulation(self, df: pd.DataFrame) -> pd.DataFrame: """Pipeline method to prepare data for clinical stratification models.""" # 1. Redact PII/PHI df = self.anonymize_cohort_data(df) - + # 2. Inject DP noise into phenotypic metrics (Age, Weight, Lab Values) - numerical_phenotypes = [c for c in df.columns if any(p in c.lower() for p in ['age', 'weight', 'height', 'level', 'value'])] + numerical_phenotypes = [ + c for c in df.columns if any(p in c.lower() for p in ["age", "weight", "height", "level", "value"]) + ] df = self.inject_differential_privacy(df, columns=numerical_phenotypes) - + return df diff --git a/compliance/samd_iec62304.py b/compliance/samd_iec62304.py index 4767bc4175..3489a8ba06 100644 --- a/compliance/samd_iec62304.py +++ b/compliance/samd_iec62304.py @@ -1,17 +1,17 @@ -import os import re -import git -import pytest from datetime import datetime -from typing import List, Dict, Any +from typing import Any + +import git from jinja2 import Template + class SaMDTraceabilityGenerator: """ Generates a Regulatory Traceability Matrix for Software as a Medical Device (SaMD). Maps Hazards (ISO 14971) to Code Changes (IEC 62304) and Verification Tests. """ - + def __init__(self, repo_path: str = "."): self.repo_path = repo_path try: @@ -19,7 +19,7 @@ def __init__(self, repo_path: str = "."): except Exception: self.repo = None - def parse_git_commits(self, limit: int = 100) -> List[Dict[str, Any]]: + def parse_git_commits(self, limit: int = 100) -> list[dict[str, Any]]: """ Scans recent commits for Hazard tags (e.g., [HAZ-001]). """ @@ -30,20 +30,22 @@ def parse_git_commits(self, limit: int = 100) -> List[Dict[str, Any]]: # Regex to find tags like [HAZ-001: Description] hazard_regex = r"\[(HAZ-\d+):\s*(.*?)\]" - for commit in self.repo.iter_commits('main', max_count=limit): + for commit in self.repo.iter_commits("main", max_count=limit): match = re.search(hazard_regex, commit.message) if match: - hazard_commits.append({ - "hazard_id": match.group(1), - "description": match.group(2), - "commit_hash": commit.hexsha, - "author": commit.author.name, - "date": datetime.fromtimestamp(commit.committed_date).isoformat() - }) - + hazard_commits.append( + { + "hazard_id": match.group(1), + "description": match.group(2), + "commit_hash": commit.hexsha, + "author": commit.author.name, + "date": datetime.fromtimestamp(commit.committed_date).isoformat(), + } + ) + return hazard_commits - def get_test_results(self, test_path: str = "tests/") -> Dict[str, str]: + def get_test_results(self, test_path: str = "tests/") -> dict[str, str]: """ Interrogates test suite to verify hazard mitigation. In production, this would read from a JUnit XML or similar test artifact. @@ -53,7 +55,7 @@ def get_test_results(self, test_path: str = "tests/") -> Dict[str, str]: "test_toxicity_gate.py": "PASSED", "test_advanced_admet.py": "PASSED", "test_retrosynthesis.py": "PASSED", - "test_safety_modules.py": "PASSED" + "test_safety_modules.py": "PASSED", } def generate_traceability_matrix(self, output_format: str = "markdown") -> str: @@ -61,8 +63,8 @@ def generate_traceability_matrix(self, output_format: str = "markdown") -> str: Correlates hazards, commits, and tests into a sterile compliance report. """ hazards = self.parse_git_commits() - tests = self.get_test_results() - + self.get_test_results() + template_str = """ # SaMD Traceability Matrix (IEC 62304 / ISO 13485) **Generated:** {{ date }} @@ -79,19 +81,18 @@ def generate_traceability_matrix(self, output_format: str = "markdown") -> str: """ template = Template(template_str) report = template.render( - hazards=hazards, - test_status="PASSED", - date=datetime.now().strftime("%Y-%m-%d %H:%M:%S") + hazards=hazards, test_status="PASSED", date=datetime.now().strftime("%Y-%m-%d %H:%M:%S") ) - + if output_format == "markdown": output_file = "compliance/traceability_matrix.md" with open(output_file, "w") as f: f.write(report) return output_file - + return report + if __name__ == "__main__": generator = SaMDTraceabilityGenerator() report_path = generator.generate_traceability_matrix() diff --git a/compliance/security/biothreat_screening.py b/compliance/security/biothreat_screening.py index 3c0eef6209..3983a01368 100644 --- a/compliance/security/biothreat_screening.py +++ b/compliance/security/biothreat_screening.py @@ -63,48 +63,45 @@ def _simulate_ache_docking(self, molecule_smiles: str) -> float: """ if not molecule_smiles: return 0.0 - + # Heuristic: estimate binding affinity from molecular properties # More atoms/complexity -> potentially better binding # More hydrophobic -> better binding to hydrophobic pocket - + mol = Chem.MolFromSmiles(molecule_smiles) if mol is None: return 0.0 - + # Base affinity (ΔG in kcal/mol, more negative = stronger binding) base_affinity = -4.0 - + # Factor 1: Molecular weight (heavier = often better for serine esterase binding) mw_factor = (mol.GetMolWt() - 100) * 0.01 # Scale: ~0 for MW 100, ~3 for MW 400 mw_factor = max(-2.0, min(2.0, mw_factor)) # Clamp to [-2, 2] - + # Factor 2: Hydrophobic atoms (Aromatic + Aliphatic carbons) - hydrophobic_atoms = sum(1 for atom in mol.GetAtoms() - if atom.GetIsAromatic() or atom.GetSymbol() == 'C') + hydrophobic_atoms = sum(1 for atom in mol.GetAtoms() if atom.GetIsAromatic() or atom.GetSymbol() == "C") hydro_factor = hydrophobic_atoms * 0.05 # ~0.5 per hydrophobic group hydro_factor = max(0, min(3.0, hydro_factor)) - + # Factor 3: Hydrogen bond donors/acceptors (serine catalytic triad has H-bond sites) - hbd = sum(1 for atom in mol.GetAtoms() - if atom.GetTotalNumHs() > 0 and atom.GetSymbol() in ['N', 'O']) - hba = sum(1 for atom in mol.GetAtoms() - if atom.GetTotalValence() > 1 and atom.GetSymbol() in ['N', 'O']) + hbd = sum(1 for atom in mol.GetAtoms() if atom.GetTotalNumHs() > 0 and atom.GetSymbol() in ["N", "O"]) + hba = sum(1 for atom in mol.GetAtoms() if atom.GetTotalValence() > 1 and atom.GetSymbol() in ["N", "O"]) hbond_factor = (hbd + hba) * 0.3 hbond_factor = max(0, min(2.0, hbond_factor)) - + # Factor 4: Rotatable bonds (more flexibility = worse, usually) - rotatable = sum(1 for bond in mol.GetBonds() - if bond.GetBondType() == Chem.BondType.SINGLE and - not bond.IsInRing()) + rotatable = sum( + 1 for bond in mol.GetBonds() if bond.GetBondType() == Chem.BondType.SINGLE and not bond.IsInRing() + ) rot_penalty = rotatable * -0.05 rot_penalty = max(-1.0, min(0, rot_penalty)) - + # Combine factors predicted_affinity = base_affinity + mw_factor + hydro_factor + hbond_factor + rot_penalty - + # Add small stochastic noise to simulate docking uncertainty (~0.5 kcal/mol) noise = Chem.RawMolDescriptors.CalcNumAtoms(mol) % 7 * 0.1 - 0.3 predicted_affinity += noise - + return float(predicted_affinity) diff --git a/configs/config.py b/configs/config.py index 6defe062cc..c231cf1e01 100644 --- a/configs/config.py +++ b/configs/config.py @@ -1,71 +1,70 @@ """ Configuration for Drug Discovery Pipeline """ -import sys # Model Configurations MODEL_CONFIGS = { - 'gnn': { - 'node_features': 8, - 'edge_features': 3, - 'hidden_dim': 128, - 'num_layers': 4, - 'num_heads': 4, - 'dropout': 0.2, - 'output_dim': 1, - 'pooling': 'attention' + "gnn": { + "node_features": 8, + "edge_features": 3, + "hidden_dim": 128, + "num_layers": 4, + "num_heads": 4, + "dropout": 0.2, + "output_dim": 1, + "pooling": "attention", }, - 'transformer': { - 'input_dim': 2048, - 'hidden_dim': 512, - 'num_layers': 6, - 'num_heads': 8, - 'dropout': 0.1, - 'output_dim': 1, + "transformer": { + "input_dim": 2048, + "hidden_dim": 512, + "num_layers": 6, + "num_heads": 8, + "dropout": 0.1, + "output_dim": 1, + }, + "mpnn": { + "node_features": 8, + "edge_features": 3, + "hidden_dim": 128, + "num_layers": 4, + "dropout": 0.2, + "output_dim": 1, }, - 'mpnn': { - 'node_features': 8, - 'edge_features': 3, - 'hidden_dim': 128, - 'num_layers': 4, - 'dropout': 0.2, - 'output_dim': 1 - } } # Training Configurations TRAINING_CONFIG = { - 'num_epochs': 200, - 'batch_size': 32, - 'learning_rate': 1e-4, - 'weight_decay': 1e-5, - 'patience': 20, - 'test_size': 0.2, - 'validation_split': 0.1 + "num_epochs": 200, + "batch_size": 32, + "learning_rate": 1e-4, + "weight_decay": 1e-5, + "patience": 20, + "test_size": 0.2, + "validation_split": 0.1, } # Data Collection Configurations DATA_CONFIG = { - 'sources': ['pubchem', 'chembl', 'approved_drugs'], - 'limit_per_source': 1000, - 'cache_dir': './data/cache', - 'min_molecular_weight': 100, - 'max_molecular_weight': 900, + "sources": ["pubchem", "chembl", "approved_drugs"], + "limit_per_source": 1000, + "cache_dir": "./data/cache", + "min_molecular_weight": 100, + "max_molecular_weight": 900, } # ADMET Thresholds ADMET_THRESHOLDS = { - 'qed_min': 0.5, - 'sa_max': 6.0, - 'lipinski_max_violations': 1, - 'molecular_weight_max': 500, - 'logp_max': 5, + "qed_min": 0.5, + "sa_max": 6.0, + "lipinski_max_violations": 1, + "molecular_weight_max": 500, + "logp_max": 5, } # Paths PATHS = { - 'cache_dir': './data/cache', - 'checkpoint_dir': './checkpoints', - 'logs_dir': './logs', - 'results_dir': './results', + "cache_dir": "./data/cache", + "checkpoint_dir": "./checkpoints", + "logs_dir": "./logs", + "results_dir": "./results", } diff --git a/cython/setup.py b/cython/setup.py index 39d2e57326..f1d7542289 100644 --- a/cython/setup.py +++ b/cython/setup.py @@ -5,9 +5,9 @@ Run: python setup.py build_ext --inplace """ -from setuptools import setup, Extension -from Cython.Build import cythonize import numpy +from Cython.Build import cythonize +from setuptools import Extension, setup ext_modules = [ Extension( diff --git a/dashboard/xai_core.py b/dashboard/xai_core.py index e0dcd87a83..10769cea94 100644 --- a/dashboard/xai_core.py +++ b/dashboard/xai_core.py @@ -81,15 +81,13 @@ # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- -_INPUT_DIM: int = 1024 # Morgan fingerprint length +_INPUT_DIM: int = 1024 # Morgan fingerprint length _HIDDEN_DIM: int = 128 _OUTPUT_DIM: int = 1 _DROPOUT_P: float = 0.15 _MC_SAMPLES: int = 10 _CONFIDENCE_THRESHOLD: float = 0.1 # variance above this triggers warning -_DEFAULT_MODEL_PATH: str = os.path.join( - os.environ.get("TMPDIR", "/tmp"), "zane_xai_surrogate.onnx" -) +_DEFAULT_MODEL_PATH: str = os.path.join(os.environ.get("TMPDIR", "/tmp"), "zane_xai_surrogate.onnx") # --------------------------------------------------------------------------- diff --git a/dashboard/xai_interface/glass_box.py b/dashboard/xai_interface/glass_box.py index 3b9a57312d..df82556a35 100644 --- a/dashboard/xai_interface/glass_box.py +++ b/dashboard/xai_interface/glass_box.py @@ -58,7 +58,7 @@ def get_3d_attention_mapping(self, molecular_features: torch.Tensor): return None molecular_features.requires_grad_() - attributions, delta = self.integrated_gradients.attribute( + attributions, _delta = self.integrated_gradients.attribute( molecular_features, target=0, return_convergence_delta=True ) return attributions diff --git a/dependency_audit.py b/dependency_audit.py index e968e0a863..eb51e3f323 100644 --- a/dependency_audit.py +++ b/dependency_audit.py @@ -4,8 +4,8 @@ import argparse import json -from importlib.util import find_spec from collections.abc import Iterable, Sequence +from importlib.util import find_spec DEFAULT_MODULES = ( "numpy", diff --git a/drug_discovery/__init__.py b/drug_discovery/__init__.py index 40f01b1f13..5500f8e48b 100644 --- a/drug_discovery/__init__.py +++ b/drug_discovery/__init__.py @@ -6,6 +6,7 @@ # Core pipeline try: from drug_discovery.pipeline import DrugDiscoveryPipeline as DrugDiscoveryPipeline + __all__.append("DrugDiscoveryPipeline") except Exception: pass @@ -13,23 +14,25 @@ # Utility modules try: from drug_discovery.data.rdkit_utils import smiles_to_sdf + from drug_discovery.docking.diffdock_placeholder import run_diffdock + from drug_discovery.docking.vina_wrapper import VinaDocker from drug_discovery.generation.torchdrug_generator import TorchDrugGenerator from drug_discovery.screening.admet_models import ADMETScreen from drug_discovery.screening.filtering import filter_admet - from drug_discovery.docking.vina_wrapper import VinaDocker - from drug_discovery.docking.diffdock_placeholder import run_diffdock from drug_discovery.structure_analysis.cif_parser import parse_cif_to_mol from drug_discovery.structure_analysis.xrpd_analysis import analyze_xrpd - __all__.extend([ - "TorchDrugGenerator", - "ADMETScreen", - "VinaDocker", - "run_diffdock", - "parse_cif_to_mol", - "analyze_xrpd", - "smiles_to_sdf", - "filter_admet" - ]) + __all__.extend( + [ + "ADMETScreen", + "TorchDrugGenerator", + "VinaDocker", + "analyze_xrpd", + "filter_admet", + "parse_cif_to_mol", + "run_diffdock", + "smiles_to_sdf", + ] + ) except Exception: pass diff --git a/drug_discovery/active_learning/__init__.py b/drug_discovery/active_learning/__init__.py index cadfec6fad..ddcd2742cd 100644 --- a/drug_discovery/active_learning/__init__.py +++ b/drug_discovery/active_learning/__init__.py @@ -38,16 +38,16 @@ from drug_discovery.active_learning.orchestrator import ActiveLearningOrchestrator as ActiveLearningOrchestrator __all__ = [ - "GaussianProcessSurrogate", - "SurrogateConfig", "AcquisitionFunction", - "ExpectedImprovement", - "UpperConfidenceBound", - "ThompsonSampling", + "ActiveLearningOrchestrator", "BayesianOptimizer", + "ExpectedImprovement", + "GaussianProcessSurrogate", "MultiFidelityOptimizer", - "ResourceAllocator", "OptimizationResult", + "ResourceAllocator", "ResourceBudget", - "ActiveLearningOrchestrator", + "SurrogateConfig", + "ThompsonSampling", + "UpperConfidenceBound", ] diff --git a/drug_discovery/active_learning/acquisition.py b/drug_discovery/active_learning/acquisition.py index 7d00b4d51b..10beeec535 100644 --- a/drug_discovery/active_learning/acquisition.py +++ b/drug_discovery/active_learning/acquisition.py @@ -137,20 +137,14 @@ def evaluate(self, X: np.ndarray) -> np.ndarray: # Use best observed if hasattr(self.surrogate, "y_buffer") and self.surrogate.y_buffer: y_obs = np.concatenate(self.surrogate.y_buffer) - if self.minimize: - target = y_obs.min() - else: - target = y_obs.max() + target = y_obs.min() if self.minimize else y_obs.max() else: target = means.mean() else: target = self.target_value # EI computation - if self.minimize: - diff = target - means - else: - diff = means - target + diff = target - means if self.minimize else means - target # Standard normal PDF and CDF z = diff / (stds + 1e-10) @@ -443,7 +437,7 @@ def evaluate(self, X: np.ndarray) -> np.ndarray: # Compute expected improvement over discretized outcomes kg = np.zeros(len(X)) - for i, (mu, sigma) in enumerate(zip(means, stds)): + for i, (mu, sigma) in enumerate(zip(means, stds, strict=False)): if sigma < 1e-10: continue @@ -452,10 +446,7 @@ def evaluate(self, X: np.ndarray) -> np.ndarray: probs = probs / (probs.sum() + 1e-10) # Best achievable with this sample - if self.minimize: - best_future = np.minimum(y_grid, mu) - else: - best_future = np.maximum(y_grid, mu) + best_future = np.minimum(y_grid, mu) if self.minimize else np.maximum(y_grid, mu) # Expected value expected_best = np.dot(probs, best_future) diff --git a/drug_discovery/active_learning/orchestrator.py b/drug_discovery/active_learning/orchestrator.py index 0b33e36d4e..9e55a025e9 100644 --- a/drug_discovery/active_learning/orchestrator.py +++ b/drug_discovery/active_learning/orchestrator.py @@ -1,9 +1,14 @@ -from typing import List, Dict, Any, Optional, Callable +from collections.abc import Callable +from typing import Any + import numpy as np import pandas as pd -from .optimizer import BayesianOptimizer + from drug_discovery.data import MolecularFeaturizer +from .optimizer import BayesianOptimizer + + class ActiveLearningOrchestrator: """ Orchestrates the Active Learning loop: @@ -11,19 +16,14 @@ class ActiveLearningOrchestrator: 2. Use Bayesian Optimization to suggest which ones to label (experiment/simulate). 3. Retrain the model on new labels. """ - - def __init__( - self, - model_trainer: Any, - featurizer: MolecularFeaturizer, - optimizer: Optional[BayesianOptimizer] = None - ): + + def __init__(self, model_trainer: Any, featurizer: MolecularFeaturizer, optimizer: BayesianOptimizer | None = None): self.trainer = model_trainer self.featurizer = featurizer self.optimizer = optimizer or BayesianOptimizer() self.labeled_data = pd.DataFrame() - - def run_cycle(self, unlabeled_smiles: List[str], oracle: Callable[[List[str]], List[float]]) -> Dict[str, Any]: + + def run_cycle(self, unlabeled_smiles: list[str], oracle: Callable[[list[str]], list[float]]) -> dict[str, Any]: """ Run one AL cycle. """ @@ -35,29 +35,29 @@ def run_cycle(self, unlabeled_smiles: List[str], oracle: Callable[[List[str]], L if fp is not None: X_pool.append(fp) valid_smiles.append(smiles) - + X_pool = np.array(X_pool) - + # 2. Suggest candidates indices = self.optimizer.suggest(X_pool) suggested_smiles = [valid_smiles[i] for i in indices] - + # 3. Query oracle (simulation or real experiment) labels = oracle(suggested_smiles) - + # 4. Update dataset - new_data = pd.DataFrame({'smiles': suggested_smiles, 'target': labels}) - self.labeled_data = pd.concat([self.labeled_data, new_data]).drop_duplicates('smiles') - + new_data = pd.DataFrame({"smiles": suggested_smiles, "target": labels}) + self.labeled_data = pd.concat([self.labeled_data, new_data]).drop_duplicates("smiles") + # 5. Tell optimizer X_new = X_pool[indices] self.optimizer.tell(X_new, np.array(labels)) - + # 6. Retrain model (optional) # self.trainer.train_on_dataframe(self.labeled_data) - + return { "num_new_labels": len(labels), "suggested_smiles": suggested_smiles, - "total_labeled": len(self.labeled_data) + "total_labeled": len(self.labeled_data), } diff --git a/drug_discovery/active_learning/uncertainty_sampler.py b/drug_discovery/active_learning/uncertainty_sampler.py index e540c35ed7..9d70f89972 100644 --- a/drug_discovery/active_learning/uncertainty_sampler.py +++ b/drug_discovery/active_learning/uncertainty_sampler.py @@ -1,8 +1,9 @@ """Active learning utilities: uncertainty-based sampling and batch selection.""" + from __future__ import annotations -from typing import List, Sequence, Tuple, Dict, Any import heapq +from collections.abc import Sequence class UncertaintySampler: @@ -16,7 +17,7 @@ class UncertaintySampler: def __init__(self, diversity_dedup: bool = True): self.diversity_dedup = diversity_dedup - def _dedupe(self, items: Sequence[str], max_keep: int) -> List[str]: + def _dedupe(self, items: Sequence[str], max_keep: int) -> list[str]: seen = set() out = [] for s in items: @@ -33,7 +34,7 @@ def select_batch( candidates: Sequence[str], uncertainties: Sequence[float], batch_size: int = 16, - ) -> List[str]: + ) -> list[str]: """Return batch of SMILES selected by highest uncertainty. candidates and uncertainties must be same length. @@ -41,7 +42,7 @@ def select_batch( if len(candidates) != len(uncertainties): raise ValueError("candidates and uncertainties must be same length") # Use a max-heap of uncertainties - heap: List[Tuple[float, int]] = [] + heap: list[tuple[float, int]] = [] for i, u in enumerate(uncertainties): heap.append((-float(u), i)) heapq.heapify(heap) @@ -56,7 +57,7 @@ def select_batch( selected = self._dedupe(selected, batch_size) return selected - def select_top_k_by_entropy(self, probs: Sequence[Sequence[float]], k: int) -> List[int]: + def select_top_k_by_entropy(self, probs: Sequence[Sequence[float]], k: int) -> list[int]: """Select top-k indices by predictive entropy. probs: list of probability vectors per candidate @@ -65,6 +66,7 @@ def select_top_k_by_entropy(self, probs: Sequence[Sequence[float]], k: int) -> L entropies = [] for p in probs: import math + e = 0.0 for pi in p: if pi > 0: diff --git a/drug_discovery/advanced/autonomous_stack.py b/drug_discovery/advanced/autonomous_stack.py index 439171c9c5..66a884951f 100644 --- a/drug_discovery/advanced/autonomous_stack.py +++ b/drug_discovery/advanced/autonomous_stack.py @@ -360,7 +360,7 @@ def __init__(self, alpha: float = 0.5): self.error_table: dict[str, float] = {} def update_errors(self, ids: Sequence[str], losses: torch.Tensor) -> None: - for sample_id, loss_val in zip(ids, losses.detach().cpu().tolist()): + for sample_id, loss_val in zip(ids, losses.detach().cpu().tolist(), strict=False): prev = self.error_table.get(sample_id, 0.0) self.error_table[sample_id] = 0.7 * prev + 0.3 * float(loss_val) @@ -464,7 +464,7 @@ def adapt( pred = cloned(support_x) loss = loss_fn(pred, support_y) grads = torch.autograd.grad(loss, cloned.parameters(), create_graph=True) - for param, grad in zip(cloned.parameters(), grads): + for param, grad in zip(cloned.parameters(), grads, strict=False): param.data = param.data - self.inner_lr * grad return cloned diff --git a/drug_discovery/agentic/__init__.py b/drug_discovery/agentic/__init__.py index 9a34dad049..3d4c629fdf 100644 --- a/drug_discovery/agentic/__init__.py +++ b/drug_discovery/agentic/__init__.py @@ -3,4 +3,4 @@ from .fda_formatter import INDGenerator from .swarm import AgenticSwarm, BioethicsAgent, TranslationAgent -__all__ = ["AgenticSwarm", "BioethicsAgent", "TranslationAgent", "INDGenerator"] +__all__ = ["AgenticSwarm", "BioethicsAgent", "INDGenerator", "TranslationAgent"] diff --git a/drug_discovery/agents/__init__.py b/drug_discovery/agents/__init__.py index b8a2a55b5b..df8a1dc205 100644 --- a/drug_discovery/agents/__init__.py +++ b/drug_discovery/agents/__init__.py @@ -5,4 +5,4 @@ from .orchestrator import AgentOrchestrator, EvaluatorAgent, GeneratorAgent, OptimizerAgent, PlannerAgent -__all__ = ["GeneratorAgent", "EvaluatorAgent", "PlannerAgent", "OptimizerAgent", "AgentOrchestrator"] +__all__ = ["AgentOrchestrator", "EvaluatorAgent", "GeneratorAgent", "OptimizerAgent", "PlannerAgent"] diff --git a/drug_discovery/ai2bmd/__init__.py b/drug_discovery/ai2bmd/__init__.py index 1aa2938280..c7db05c9bc 100644 --- a/drug_discovery/ai2bmd/__init__.py +++ b/drug_discovery/ai2bmd/__init__.py @@ -3,4 +3,4 @@ Advanced MD with equivariant diffusion. """ -from .ai2bmd_dynamics import AI2BMDDynamics, DynamicsResult \ No newline at end of file +from .ai2bmd_dynamics import AI2BMDDynamics, DynamicsResult diff --git a/drug_discovery/ai2bmd/ai2bmd_dynamics.py b/drug_discovery/ai2bmd/ai2bmd_dynamics.py index 8356ebea97..dec1ffe80e 100644 --- a/drug_discovery/ai2bmd/ai2bmd_dynamics.py +++ b/drug_discovery/ai2bmd/ai2bmd_dynamics.py @@ -1,15 +1,16 @@ from __future__ import annotations -import asyncio +from collections.abc import Sequence from dataclasses import dataclass -from typing import Any, Sequence try: import ray + _RAY_AVAILABLE = True except ImportError: _RAY_AVAILABLE = False + @dataclass class DynamicsResult: smiles: str @@ -19,9 +20,10 @@ class DynamicsResult: converged: bool = False success: bool = False + class AI2BMDDynamics: """AI2BMD proxy for fast biomolecular dynamics (2025). - + Uses diffusion models for long-timescale MD. """ @@ -35,34 +37,29 @@ async def simulate_batch(self, complexes: Sequence[tuple[str, str]]) -> list[Dyn # Heuristic stability RMSD based on SMILES complexity # More complex ligands may have higher RMSD smiles_len = len(smiles) - heavy_atoms = smiles.count('C') + smiles.count('N') + smiles.count('O') + smiles.count('S') - + heavy_atoms = smiles.count("C") + smiles.count("N") + smiles.count("O") + smiles.count("S") + # RMSD typically ranges 1.0 - 4.0 Angstroms base_rmsd = 1.5 complexity_factor = min(2.0, heavy_atoms * 0.1) rmsd = base_rmsd + complexity_factor + (smiles_len % 5) * 0.1 - + # Binding affinity (ΔG) based on interaction features # More hydrophobic atoms -> better binding - hydrophobic_ratio = smiles.count('C') / max(1, smiles_len * 0.5) - + hydrophobic_ratio = smiles.count("C") / max(1, smiles_len * 0.5) + # Base ΔG ranges -6 to -12 kcal/mol for known binders base_delta = -7.0 affinity_factor = -hydrophobic_ratio * 2.0 delta = max(-12.0, min(-4.0, base_delta + affinity_factor)) - + # PDB complexity affects convergence - pdb_atoms = pdb.count('ATOM') + pdb_atoms = pdb.count("ATOM") converged = rmsd < 3.5 or pdb_atoms > 1000 success = rmsd < 4.0 - + res = DynamicsResult( - smiles, - pdb, - stability_rmsd=rmsd, - binding_delta=delta, - converged=converged, - success=success + smiles, pdb, stability_rmsd=rmsd, binding_delta=delta, converged=converged, success=success ) results.append(res) - return results \ No newline at end of file + return results diff --git a/drug_discovery/alphafold3/__init__.py b/drug_discovery/alphafold3/__init__.py index 7ba3021877..b63c49c3ea 100644 --- a/drug_discovery/alphafold3/__init__.py +++ b/drug_discovery/alphafold3/__init__.py @@ -3,6 +3,6 @@ Uses DiffDock + OpenFold for structure/pocket, with Ray-distributed batching. """ -from .alphafold3_docking import AlphaFold3Docking, AF3Result +from .alphafold3_docking import AF3Result, AlphaFold3Docking -__all__ = ["AlphaFold3Docking", "AF3Result"] \ No newline at end of file +__all__ = ["AF3Result", "AlphaFold3Docking"] diff --git a/drug_discovery/alphafold3/alphafold3_docking.py b/drug_discovery/alphafold3/alphafold3_docking.py index 8e1e668d4c..3b2b7e84e8 100644 --- a/drug_discovery/alphafold3/alphafold3_docking.py +++ b/drug_discovery/alphafold3/alphafold3_docking.py @@ -12,11 +12,13 @@ try: import ray + _RAY_AVAILABLE = True except ImportError: ray = None _RAY_AVAILABLE = False + @dataclass class AF3Result: smiles: str @@ -38,9 +40,10 @@ def as_dict(self) -> dict[str, Any]: "error": self.error, } + class AlphaFold3Docking: """Proxy for AlphaFold3 ligand binding using DiffDock/OpenFold stack. - + Falls back to RDKit pocket estimation if heavy deps unavailable. """ @@ -75,10 +78,11 @@ async def _dock_ray(self, smiles_list: Sequence[str]) -> list[AF3Result]: def dock_task(smiles: str): # Delegate to DiffDock if available from drug_discovery.physics import DiffDockAdapter - adapter = DiffDockAdapter() + + DiffDockAdapter() # Mock return AF3Result(smiles, "mock", ["pocket"], 0.8).as_dict() futures = [dock_task.remote(smi) for smi in smiles_list] raw = await asyncio.gather(*[asyncio.wrap_future(f.get()) for f in futures]) - return [AF3Result(**r) for r in raw] \ No newline at end of file + return [AF3Result(**r) for r in raw] diff --git a/drug_discovery/apex_orchestrator.py b/drug_discovery/apex_orchestrator.py index ff0bb7ccac..06722df9fc 100644 --- a/drug_discovery/apex_orchestrator.py +++ b/drug_discovery/apex_orchestrator.py @@ -77,9 +77,7 @@ async def _run_fep_scoring(self, context: dict[str, Any]) -> dict[str, Any]: try: results = await self.physics_oracle.score_batch(smiles_list) scores = [ - {"smiles": r.smiles, "delta_g": r.delta_g, "converged": r.converged} - for r in results - if r.success + {"smiles": r.smiles, "delta_g": r.delta_g, "converged": r.converged} for r in results if r.success ] logger.info("FEP scoring complete: %d/%d successful", len(scores), len(smiles_list)) return {"fep_scores": scores, "skipped": False} diff --git a/drug_discovery/biomarker_discovery/ml_discovery.py b/drug_discovery/biomarker_discovery/ml_discovery.py index 243f6f21d2..59fc480a9d 100644 --- a/drug_discovery/biomarker_discovery/ml_discovery.py +++ b/drug_discovery/biomarker_discovery/ml_discovery.py @@ -1,63 +1,58 @@ -import pandas as pd import numpy as np +import pandas as pd from sklearn.ensemble import RandomForestClassifier +from sklearn.metrics import roc_auc_score from sklearn.model_selection import train_test_split -from sklearn.metrics import roc_auc_score, precision_recall_curve -from typing import List, Dict, Any, Tuple + class BiomarkerMLDiscovery: """ Machine Learning methods for biomarker identification and ranking. """ - + def __init__(self, data: pd.DataFrame, target_col: str): self.data = data self.target_col = target_col - - def rank_features_by_importance(self, features: List[str]) -> pd.DataFrame: + + def rank_features_by_importance(self, features: list[str]) -> pd.DataFrame: """ Use Random Forest to rank features by their importance in predicting the target. """ X = self.data[features] y = self.data[self.target_col] - + # Simple binary classification or regression check if y.dtype == object or len(np.unique(y)) < 10: model = RandomForestClassifier(n_estimators=100, random_state=42) else: from sklearn.ensemble import RandomForestRegressor + model = RandomForestRegressor(n_estimators=100, random_state=42) - + model.fit(X, y) - + importances = model.feature_importances_ indices = np.argsort(importances)[::-1] - + ranked_features = [] for f in range(X.shape[1]): - ranked_features.append({ - "feature": features[indices[f]], - "importance": importances[indices[f]] - }) - + ranked_features.append({"feature": features[indices[f]], "importance": importances[indices[f]]}) + return pd.DataFrame(ranked_features) - def evaluate_biomarker_panel(self, panel_features: List[str]) -> Dict[str, float]: + def evaluate_biomarker_panel(self, panel_features: list[str]) -> dict[str, float]: """ Evaluate the predictive power of a set of biomarkers. """ X = self.data[panel_features] y = self.data[self.target_col] - + X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) - + model = RandomForestClassifier(n_estimators=100, random_state=42) model.fit(X_train, y_train) - + probs = model.predict_proba(X_test)[:, 1] auc = roc_auc_score(y_test, probs) - - return { - "auc_roc": auc, - "num_features": len(panel_features) - } + + return {"auc_roc": auc, "num_features": len(panel_features)} diff --git a/drug_discovery/biomarker_discovery/statistical_analysis.py b/drug_discovery/biomarker_discovery/statistical_analysis.py index c92d2296d1..9b27009df6 100644 --- a/drug_discovery/biomarker_discovery/statistical_analysis.py +++ b/drug_discovery/biomarker_discovery/statistical_analysis.py @@ -1,13 +1,13 @@ import numpy as np import pandas as pd from scipy import stats -from typing import Dict, List, Optional, Tuple + class BiomarkerStatisticalAnalysis: """ Statistical methods for biomarker identification. """ - + def __init__(self, data: pd.DataFrame, group_col: str): """ Args: @@ -20,40 +20,42 @@ def __init__(self, data: pd.DataFrame, group_col: str): if len(self.groups) < 2: raise ValueError("Data must contain at least two groups for comparison.") - def differential_expression(self, features: List[str]) -> pd.DataFrame: + def differential_expression(self, features: list[str]) -> pd.DataFrame: """ Perform t-tests for each feature between two groups. - + Returns: DataFrame with t-statistic, p-value, and fold change. """ results = [] group1 = self.data[self.data[self.group_col] == self.groups[0]] group2 = self.data[self.data[self.group_col] == self.groups[1]] - + for feature in features: if feature == self.group_col: continue - + v1 = group1[feature].dropna() v2 = group2[feature].dropna() - + if len(v1) < 2 or len(v2) < 2: continue - + t_stat, p_val = stats.ttest_ind(v1, v2) fold_change = v1.mean() / (v2.mean() + 1e-9) - - results.append({ - "feature": feature, - "t_statistic": t_stat, - "p_value": p_val, - "log2_fold_change": np.log2(fold_change + 1e-9) - }) - + + results.append( + { + "feature": feature, + "t_statistic": t_stat, + "p_value": p_val, + "log2_fold_change": np.log2(fold_change + 1e-9), + } + ) + return pd.DataFrame(results).sort_values("p_value") - def correlation_analysis(self, target_feature: str, features: List[str]) -> pd.DataFrame: + def correlation_analysis(self, target_feature: str, features: list[str]) -> pd.DataFrame: """ Calculate correlation between features and a target feature. """ @@ -61,12 +63,8 @@ def correlation_analysis(self, target_feature: str, features: List[str]) -> pd.D for feature in features: if feature == target_feature: continue - + corr, p_val = stats.pearsonr(self.data[feature], self.data[target_feature]) - correlations.append({ - "feature": feature, - "correlation": corr, - "p_value": p_val - }) - + correlations.append({"feature": feature, "correlation": corr, "p_value": p_val}) + return pd.DataFrame(correlations).sort_values("p_value") diff --git a/drug_discovery/causal_discovery/causal_graph.py b/drug_discovery/causal_discovery/causal_graph.py index 336fdbcc77..130924dd7b 100644 --- a/drug_discovery/causal_discovery/causal_graph.py +++ b/drug_discovery/causal_discovery/causal_graph.py @@ -1,20 +1,20 @@ import networkx as nx import pandas as pd -from typing import List, Dict, Any, Tuple + class CausalGraph: """ Representation and discovery of causal graphs. """ - - def __init__(self, nodes: List[str] = None): + + def __init__(self, nodes: list[str] | None = None): self.graph = nx.DiGraph() if nodes: self.graph.add_nodes_from(nodes) - + def add_causal_link(self, source: str, target: str, strength: float = 1.0): self.graph.add_edge(source, target, weight=strength) - + def discover_from_data(self, data: pd.DataFrame, method: str = "correlation_threshold"): """ Discover causal structure from data (simplified). @@ -27,9 +27,9 @@ def discover_from_data(self, data: pd.DataFrame, method: str = "correlation_thre # In a real scenario, we'd use PC algorithm or similar # Here we just add a directed edge based on column order as placeholder self.add_causal_link(corr.columns[i], corr.columns[j], strength=corr.iloc[i, j]) - - def get_descendants(self, node: str) -> List[str]: + + def get_descendants(self, node: str) -> list[str]: return list(nx.descendants(self.graph, node)) - - def get_ancestors(self, node: str) -> List[str]: + + def get_ancestors(self, node: str) -> list[str]: return list(nx.ancestors(self.graph, node)) diff --git a/drug_discovery/causal_discovery/inference.py b/drug_discovery/causal_discovery/inference.py index 12818aa495..11ca18d692 100644 --- a/drug_discovery/causal_discovery/inference.py +++ b/drug_discovery/causal_discovery/inference.py @@ -1,27 +1,26 @@ import pandas as pd -import numpy as np from sklearn.linear_model import LinearRegression -from typing import Dict, List, Any + class CausalInference: """ Perform causal inference tasks like effect estimation. """ - + def __init__(self, data: pd.DataFrame): self.data = data - - def estimate_treatment_effect(self, treatment: str, outcome: str, confounders: List[str]) -> float: + + def estimate_treatment_effect(self, treatment: str, outcome: str, confounders: list[str]) -> float: """ Estimate the effect of a treatment on an outcome, adjusting for confounders. Uses a simple adjustment formula (backdoor adjustment via linear regression). """ - X = self.data[[treatment] + confounders] + X = self.data[[treatment, *confounders]] y = self.data[outcome] - + model = LinearRegression() model.fit(X, y) - + # The coefficient of the treatment is the estimated average treatment effect (ATE) # under the assumption of linearity and no unmeasured confounders. ate = model.coef_[0] @@ -32,4 +31,4 @@ def counterfactual_prediction(self, individual_data: pd.Series, treatment: str, Predict what would happen if the treatment was set to a new value. """ # Simplified implementation - return 0.0 # Placeholder + return 0.0 # Placeholder diff --git a/drug_discovery/commercial/fda_drug_matcher.py b/drug_discovery/commercial/fda_drug_matcher.py index fc5acae96a..940059c81a 100644 --- a/drug_discovery/commercial/fda_drug_matcher.py +++ b/drug_discovery/commercial/fda_drug_matcher.py @@ -1,14 +1,15 @@ +import logging + import pandas as pd from rdkit import Chem, DataStructs from rdkit.Chem import AllChem -from typing import Dict, Optional, List -import logging logger = logging.getLogger(__name__) + class CommercialDrugMapper: def __init__(self): - self.fda_db: Optional[pd.DataFrame] = None + self.fda_db: pd.DataFrame | None = None self.fingerprints = [] def load_fda_orange_book(self, database_path: str): @@ -23,16 +24,26 @@ def load_fda_orange_book(self, database_path: str): except Exception as e: logger.error(f"Error loading FDA Orange Book: {e}") # Fallback for demonstration if file doesn't exist - self.fda_db = pd.DataFrame([ - {"drug_name": "Imatinib", "smiles": "CC1=C(C=C(C=C1)NC(=O)C2=CC=C(C=C2)CN3CCN(CC3)C)NC4=NC=CC(=N4)C5=CN=CC=C5", "commercial_dose": "400mg daily"}, - {"drug_name": "Metformin", "smiles": "CN(C)C(=N)N=C(N)N", "commercial_dose": "500mg twice daily"}, - {"drug_name": "Atorvastatin", "smiles": "CC(C)C1=C(C(=C(N1CC[C@H](C[C@H](CC(=O)O)O)O)C2=CC=C(C=C2)F)C3=CC=CC=C3)C(=O)NC4=CC=CC=C4", "commercial_dose": "20mg daily"} - ]) + self.fda_db = pd.DataFrame( + [ + { + "drug_name": "Imatinib", + "smiles": "CC1=C(C=C(C=C1)NC(=O)C2=CC=C(C=C2)CN3CCN(CC3)C)NC4=NC=CC(=N4)C5=CN=CC=C5", + "commercial_dose": "400mg daily", + }, + {"drug_name": "Metformin", "smiles": "CN(C)C(=N)N=C(N)N", "commercial_dose": "500mg twice daily"}, + { + "drug_name": "Atorvastatin", + "smiles": "CC(C)C1=C(C(=C(N1CC[C@H](C[C@H](CC(=O)O)O)O)C2=CC=C(C=C2)F)C3=CC=CC=C3)C(=O)NC4=CC=CC=C4", + "commercial_dose": "20mg daily", + }, + ] + ) self._precompute_fingerprints() def _precompute_fingerprints(self): self.fingerprints = [] - for smiles in self.fda_db['smiles']: + for smiles in self.fda_db["smiles"]: mol = Chem.MolFromSmiles(smiles) if mol: fp = AllChem.GetMorganFingerprintAsBitVect(mol, 2, nBits=2048) @@ -42,7 +53,7 @@ def _precompute_fingerprints(self): def find_closest_commercial_match(self, generated_smiles: str) -> dict: """ - Calculates Morgan Fingerprint Tanimoto similarity between ZANE's de novo molecule + Calculates Morgan Fingerprint Tanimoto similarity between ZANE's de novo molecule and all FDA-approved drugs. """ if self.fda_db is None or not self.fingerprints: @@ -53,42 +64,38 @@ def find_closest_commercial_match(self, generated_smiles: str) -> dict: return {"closest_drug": "Invalid SMILES", "similarity": 0.0, "commercial_dose": "N/A", "smiles": ""} query_fp = AllChem.GetMorganFingerprintAsBitVect(query_mol, 2, nBits=2048) - + max_sim = -1.0 best_match_idx = -1 - + for i, fp in enumerate(self.fingerprints): if fp: sim = DataStructs.TanimotoSimilarity(query_fp, fp) if sim > max_sim: max_sim = sim best_match_idx = i - + if best_match_idx != -1: match = self.fda_db.iloc[best_match_idx] return { - "closest_drug": match['drug_name'], + "closest_drug": match["drug_name"], "similarity": round(max_sim, 4), - "commercial_dose": match['commercial_dose'], - "smiles": match['smiles'] + "commercial_dose": match["commercial_dose"], + "smiles": match["smiles"], } - + return {"closest_drug": "No Match", "similarity": 0.0, "commercial_dose": "N/A", "smiles": ""} - def compare_compounds(self, zane_compounds: List[dict], commercial_match: dict) -> dict: + def compare_compounds(self, zane_compounds: list[dict], commercial_match: dict) -> dict: """ Compares ZANE's multi-compound drug with the commercial equivalent. """ - zane_smiles_set = {c['smiles'] for c in zane_compounds} - comm_smiles = commercial_match.get('smiles', '') - + comm_smiles = commercial_match.get("smiles", "") + # Extra: In ZANE but not the main commercial ingredient - extra = [c['smiles'] for c in zane_compounds if c['smiles'] != comm_smiles] - + extra = [c["smiles"] for c in zane_compounds if c["smiles"] != comm_smiles] + # Missing: In commercial but not in ZANE (Mocked for demonstration) missing = ["Magnesium Stearate (Excipient)", "Hypromellose (Coating)"] if comm_smiles else [] - - return { - "extra_compounds": extra[:5], # Limit display - "missing_compounds": missing - } + + return {"extra_compounds": extra[:5], "missing_compounds": missing} # Limit display diff --git a/drug_discovery/compliance/__init__.py b/drug_discovery/compliance/__init__.py index 46d67047e4..aae55069fa 100644 --- a/drug_discovery/compliance/__init__.py +++ b/drug_discovery/compliance/__init__.py @@ -17,7 +17,7 @@ compliance_log as compliance_log, ) - __all__.extend(["AuditLedger", "AuditEntry", "compliance_log"]) + __all__.extend(["AuditEntry", "AuditLedger", "compliance_log"]) except ImportError: pass @@ -41,6 +41,6 @@ require_signature as require_signature, ) - __all__.extend(["RBACManager", "User", "Role", "Permission", "require_permission", "require_signature"]) + __all__.extend(["Permission", "RBACManager", "Role", "User", "require_permission", "require_signature"]) except ImportError: pass diff --git a/drug_discovery/compliance/audit_trail.py b/drug_discovery/compliance/audit_trail.py index bc0bb28b5c..0a31e812c8 100644 --- a/drug_discovery/compliance/audit_trail.py +++ b/drug_discovery/compliance/audit_trail.py @@ -12,10 +12,10 @@ import hashlib import json import logging -from dataclasses import dataclass, field, asdict +from dataclasses import dataclass, field from datetime import datetime, timezone from enum import Enum -from typing import Any, Optional +from typing import Any logger = logging.getLogger(__name__) @@ -37,7 +37,7 @@ class AuditEventType(Enum): @dataclass class ComplianceAuditEntry: """Single audit trail entry (immutable). - + Follows FDA Part 11 ALCOA+ principles: - Attributable: Who performed the action - Legible: Readable text @@ -50,25 +50,25 @@ class ComplianceAuditEntry: event_type: AuditEventType timestamp: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) user_id: str = "system" - + # Event details - compound_id: Optional[str] = None - smiles: Optional[str] = None + compound_id: str | None = None + smiles: str | None = None details: dict[str, Any] = field(default_factory=dict) - + # Integrity verification event_hash: str = "" # SHA256 of event data - previous_hash: Optional[str] = None # Hash chaining for integrity - + previous_hash: str | None = None # Hash chaining for integrity + # Regulatory metadata audit_id: str = "" # Unique identifier - - def compute_hash(self, previous_hash: Optional[str] = None) -> str: + + def compute_hash(self, previous_hash: str | None = None) -> str: """Compute cryptographic hash for immutability verification. - + Args: previous_hash: Hash of previous entry for chain verification - + Returns: SHA256 hash of this entry """ @@ -82,7 +82,7 @@ def compute_hash(self, previous_hash: Optional[str] = None) -> str: "details": self.details, "previous_hash": previous_hash, } - + json_str = json.dumps(data, sort_keys=True, separators=(",", ":")) return hashlib.sha256(json_str.encode()).hexdigest() @@ -90,31 +90,31 @@ def compute_hash(self, previous_hash: Optional[str] = None) -> str: @dataclass class AuditTrail: """Immutable audit trail for compliance. - + Maintains hash chain to detect any modification of entries. """ entries: list[ComplianceAuditEntry] = field(default_factory=list) chain_verified: bool = True - last_verification: Optional[datetime] = None + last_verification: datetime | None = None def add_entry( self, event_type: AuditEventType, user_id: str = "system", - compound_id: Optional[str] = None, - smiles: Optional[str] = None, - details: Optional[dict[str, Any]] = None, + compound_id: str | None = None, + smiles: str | None = None, + details: dict[str, Any] | None = None, ) -> ComplianceAuditEntry: """Add new entry to audit trail. - + Args: event_type: Type of event user_id: User performing action compound_id: Optional compound identifier smiles: Optional SMILES string details: Additional event details - + Returns: The created audit entry """ @@ -134,10 +134,7 @@ def add_entry( entry.audit_id = self._generate_audit_id(entry) self.entries.append(entry) - logger.info( - f"Audit entry recorded: {event_type.value} " - f"(audit_id={entry.audit_id})" - ) + logger.info(f"Audit entry recorded: {event_type.value} " f"(audit_id={entry.audit_id})") return entry @@ -148,7 +145,7 @@ def _generate_audit_id(self, entry: ComplianceAuditEntry) -> str: def verify_chain_integrity(self) -> bool: """Verify entire audit chain hasn't been tampered with. - + Returns: True if chain is valid, False if tampering detected """ @@ -188,7 +185,7 @@ def get_entries_for_compound(self, compound_id: str) -> list[ComplianceAuditEntr def get_entries_since( self, start_time: datetime, - event_type: Optional[AuditEventType] = None, + event_type: AuditEventType | None = None, ) -> list[ComplianceAuditEntry]: """Get audit entries since a specific time.""" entries = [e for e in self.entries if e.timestamp >= start_time] @@ -239,14 +236,14 @@ def _serialize_entry(self, entry: ComplianceAuditEntry) -> dict[str, Any]: class ComplianceAuditLogger: """Convenience logger for recording compliance events.""" - def __init__(self, audit_trail: Optional[AuditTrail] = None): + def __init__(self, audit_trail: AuditTrail | None = None): """Initialize audit logger.""" self.audit_trail = audit_trail or AuditTrail() def log_compound_screened( self, smiles: str, - compound_id: Optional[str] = None, + compound_id: str | None = None, user_id: str = "system", ) -> ComplianceAuditEntry: """Log compound screening event.""" @@ -303,11 +300,7 @@ def log_approval_decision( user_id: str = "system", ) -> ComplianceAuditEntry: """Log approval decision.""" - event_type = ( - AuditEventType.APPROVAL_DECISION - if decision == "approved" - else AuditEventType.REJECTION_DECISION - ) + event_type = AuditEventType.APPROVAL_DECISION if decision == "approved" else AuditEventType.REJECTION_DECISION return self.audit_trail.add_entry( event_type, user_id=user_id, diff --git a/drug_discovery/compliance/rbac.py b/drug_discovery/compliance/rbac.py index 14c2cf19c1..8453e8bf80 100644 --- a/drug_discovery/compliance/rbac.py +++ b/drug_discovery/compliance/rbac.py @@ -189,8 +189,7 @@ def check_permission(self, user: User, permission: Permission) -> None: """Raise :class:`PermissionError` if the user lacks *permission*.""" if not user.has_permission(permission): raise PermissionError( - f"User {user.user_id} (role={user.role.name}) " - f"lacks permission {permission.value}" + f"User {user.user_id} (role={user.role.name}) " f"lacks permission {permission.value}" ) def verify_signature( @@ -212,9 +211,7 @@ def verify_signature( if not user.has_permission(Permission.SIGN_PREDICTION) and not user.has_permission( Permission.APPROVE_CHECKPOINT ): - raise SignatureError( - f"User {user.user_id} lacks signing permission" - ) + raise SignatureError(f"User {user.user_id} lacks signing permission") user.last_auth = time.monotonic() @@ -224,9 +221,7 @@ def verify_signature( "role": user.role.name, "reason": reason, "timestamp": time.time(), - "signature_hash": hashlib.sha256( - f"{user.user_id}:{reason}:{time.time()}".encode() - ).hexdigest(), + "signature_hash": hashlib.sha256(f"{user.user_id}:{reason}:{time.time()}".encode()).hexdigest(), } logger.info("Electronic signature recorded for user %s: %s", user.user_id, reason) return signature_record @@ -269,9 +264,7 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: if rbac: rbac.check_permission(user, permission) elif not user.has_permission(permission): - raise PermissionError( - f"User {user.user_id} lacks {permission.value}" - ) + raise PermissionError(f"User {user.user_id} lacks {permission.value}") return fn(*args, **kwargs) return wrapper diff --git a/drug_discovery/continuous_improvement/__init__.py b/drug_discovery/continuous_improvement/__init__.py index ba1e10fa9b..055fee538c 100644 --- a/drug_discovery/continuous_improvement/__init__.py +++ b/drug_discovery/continuous_improvement/__init__.py @@ -19,10 +19,10 @@ ) __all__ = [ + "ConceptDriftDetector", "ContinuousImprovementSystem", "DataDriftDetector", - "ConceptDriftDetector", - "PerformanceMonitor", "DriftReport", "PerformanceMetric", + "PerformanceMonitor", ] diff --git a/drug_discovery/dashboard.py b/drug_discovery/dashboard.py index 97c05a81c1..4954cc20f0 100644 --- a/drug_discovery/dashboard.py +++ b/drug_discovery/dashboard.py @@ -882,7 +882,7 @@ def _compute_combo_rankings( def _metric_bar(value: float, min_value: float, max_value: float, width: int = 20) -> str: span = max(max_value - min_value, 1e-9) ratio = max(0.0, min(1.0, (value - min_value) / span)) - filled = int(round(ratio * width)) + filled = round(ratio * width) return "[" + ("#" * filled) + ("-" * (width - filled)) + "]" diff --git a/drug_discovery/data/__init__.py b/drug_discovery/data/__init__.py index a6c4492808..a9af4113e8 100644 --- a/drug_discovery/data/__init__.py +++ b/drug_discovery/data/__init__.py @@ -4,9 +4,17 @@ from drug_discovery.data.collector import DataCollector as DataCollector from drug_discovery.data.dataset import ( MolecularDataset as MolecularDataset, + ) + from drug_discovery.data.dataset import ( MolecularFeaturizer as MolecularFeaturizer, + ) + from drug_discovery.data.dataset import ( murcko_scaffold_kfold_split_molecular as murcko_scaffold_kfold_split_molecular, + ) + from drug_discovery.data.dataset import ( murcko_scaffold_split_molecular as murcko_scaffold_split_molecular, + ) + from drug_discovery.data.dataset import ( train_test_split_molecular as train_test_split_molecular, ) @@ -14,8 +22,8 @@ "DataCollector", "MolecularDataset", "MolecularFeaturizer", - "train_test_split_molecular", "murcko_scaffold_split_molecular", + "train_test_split_molecular", ] except ImportError: # Keep data module lazy when dependencies are unavailable. diff --git a/drug_discovery/data/collector.py b/drug_discovery/data/collector.py index 7c5d27f91b..0df4f10d0e 100644 --- a/drug_discovery/data/collector.py +++ b/drug_discovery/data/collector.py @@ -3,14 +3,17 @@ from __future__ import annotations import logging -import time import os +import time +from collections.abc import Callable from pathlib import Path -from typing import Any, Callable, TypeVar, cast +from typing import Any, TypeVar, cast import pandas as pd + try: from pymongo import MongoClient + _PYMONGO = True except ImportError: _PYMONGO = False @@ -24,11 +27,13 @@ class DataCollector: """Collect and merge records from multiple sources with safe fallbacks.""" - def __init__(self, cache_dir: str = "./data/cache", api_keys: dict[str, str] | None = None, use_mongodb: bool = True): + def __init__( + self, cache_dir: str = "./data/cache", api_keys: dict[str, str] | None = None, use_mongodb: bool = True + ): self.cache_dir = cache_dir self.api_keys = api_keys or {} Path(cache_dir).mkdir(parents=True, exist_ok=True) - + self.use_mongodb = use_mongodb and _PYMONGO if self.use_mongodb and MongoClient: mongodb_uri = os.getenv("MONGODB_URI", "mongodb://localhost:27017") @@ -109,7 +114,7 @@ def _set_cache(self, source: str, query: str, df: pd.DataFrame) -> None: self.collection.update_one( {"source": source, "query": query}, {"$set": {"source": source, "query": query, "data": data, "timestamp": time.time()}}, - upsert=True + upsert=True, ) except Exception as e: logger.warning(f"Failed to write to MongoDB cache: {e}") @@ -151,7 +156,7 @@ def _fetch() -> list[Any]: return out mask = out["smiles"].map(self._is_valid_smiles).astype(bool) filtered = cast(pd.DataFrame, out.loc[mask].copy().reset_index(drop=True)) - + self._set_cache("pubchem", f"{query}_{namespace}_{limit}", filtered) return filtered @@ -183,7 +188,9 @@ def _fetch() -> Any: query = query.filter(target_chembl_id=target_results[0].get("target_chembl_id")) if activity_type: query = query.filter(standard_type=activity_type) - return query.only(["molecule_chembl_id", "canonical_smiles", "molecule_pref_name"])[: max(0, int(limit))] + return query.only(["molecule_chembl_id", "canonical_smiles", "molecule_pref_name"])[ + : max(0, int(limit)) + ] return new_client.molecule.filter(molecule_structures__isnull=False).only( ["molecule_chembl_id", "molecule_structures", "pref_name"] )[: max(0, int(limit))] @@ -200,7 +207,12 @@ def _fetch() -> Any: rows.append( { "smiles": str(smiles), - "name": str(item.get("pref_name") or item.get("molecule_pref_name") or item.get("molecule_chembl_id") or ""), + "name": str( + item.get("pref_name") + or item.get("molecule_pref_name") + or item.get("molecule_chembl_id") + or "" + ), "source": "chembl", } ) @@ -213,7 +225,7 @@ def _fetch() -> Any: return out mask = out["smiles"].map(self._is_valid_smiles).astype(bool) filtered = cast(pd.DataFrame, out.loc[mask].copy().reset_index(drop=True)) - + self._set_cache("chembl", query_key, filtered) return filtered @@ -299,7 +311,9 @@ def collect_from_pdb( } response = requests.post(search_url, json=payload, timeout=20) response.raise_for_status() - ids = [item.get("identifier") for item in response.json().get("result_set", []) if item.get("identifier")] + ids = [ + item.get("identifier") for item in response.json().get("result_set", []) if item.get("identifier") + ] except Exception as exc: logger.warning("PDB search failed: %s", exc) return pd.DataFrame() diff --git a/drug_discovery/data/dataset.py b/drug_discovery/data/dataset.py index 0ee6fd18cd..b9127c2444 100644 --- a/drug_discovery/data/dataset.py +++ b/drug_discovery/data/dataset.py @@ -2,8 +2,8 @@ from __future__ import annotations -from dataclasses import dataclass import random +from dataclasses import dataclass import torch from torch.utils.data import Dataset, Subset @@ -313,7 +313,7 @@ def train_test_split_molecular(dataset: MolecularDataset, test_size: float = 0.2 if n == 0: return Subset(dataset, []), Subset(dataset, []) - test_n = int(round(n * float(test_size))) + test_n = round(n * float(test_size)) if n > 1: test_n = max(1, test_n) test_n = min(test_n, n) @@ -359,7 +359,7 @@ def murcko_scaffold_split_molecular(dataset: MolecularDataset, test_size: float rng.shuffle(scaffold_groups) scaffold_groups.sort(key=len, reverse=True) - target_test_n = min(n, max(1 if n > 1 else 0, int(round(n * float(test_size))))) + target_test_n = min(n, max(1 if n > 1 else 0, round(n * float(test_size)))) test_indices: list[int] = [] train_indices: list[int] = [] diff --git a/drug_discovery/data/feature_store.py b/drug_discovery/data/feature_store.py index d829bb1c6f..43c5db62f9 100644 --- a/drug_discovery/data/feature_store.py +++ b/drug_discovery/data/feature_store.py @@ -6,14 +6,16 @@ """ import logging -import pickle import os +import pickle from pathlib import Path -from typing import Dict, Optional, Any, List +from typing import Any + import numpy as np -import torch + try: from pymongo import MongoClient + _PYMONGO = True except ImportError: _PYMONGO = False @@ -37,8 +39,8 @@ def __init__(self, store_path: str = "./cache/features", use_mongodb: bool = Tru self.store_path.mkdir(parents=True, exist_ok=True) # In-memory cache - self._cache: Dict[str, Any] = {} - + self._cache: dict[str, Any] = {} + self.use_mongodb = use_mongodb and _PYMONGO if self.use_mongodb and MongoClient: mongodb_uri = os.getenv("MONGODB_URI", "mongodb://localhost:27017") @@ -57,7 +59,7 @@ def store_embedding( key: str, embedding: np.ndarray, feature_type: str = "molecule", - metadata: Optional[Dict] = None, + metadata: dict | None = None, ) -> None: """ Store an embedding vector. @@ -77,11 +79,7 @@ def store_embedding( if self.use_mongodb: try: - self.collection.update_one( - {"key": key, "feature_type": feature_type}, - {"$set": data}, - upsert=True - ) + self.collection.update_one({"key": key, "feature_type": feature_type}, {"$set": data}, upsert=True) except Exception as e: logger.error(f"Failed to store embedding in MongoDB for {key}: {e}") else: @@ -99,7 +97,7 @@ def retrieve_embedding( self, key: str, feature_type: str = "molecule", - ) -> Optional[np.ndarray]: + ) -> np.ndarray | None: """ Retrieve an embedding vector. @@ -140,15 +138,15 @@ def retrieve_embedding( return emb except Exception as e: logger.error(f"Failed to load embedding for {key}: {e}") - + return None def store_batch( self, - keys: List[str], + keys: list[str], embeddings: np.ndarray, feature_type: str = "molecule", - metadata_list: Optional[List[Dict]] = None, + metadata_list: list[dict] | None = None, ) -> None: """ Store multiple embeddings efficiently. @@ -165,29 +163,26 @@ def store_batch( if self.use_mongodb: operations = [] from pymongo import UpdateOne - for key, embedding, metadata in zip(keys, embeddings, metadata_list): + + for key, embedding, metadata in zip(keys, embeddings, metadata_list, strict=False): data = { "key": key, "feature_type": feature_type, "embedding": embedding.tolist() if isinstance(embedding, np.ndarray) else embedding, "metadata": metadata or {}, } - operations.append(UpdateOne( - {"key": key, "feature_type": feature_type}, - {"$set": data}, - upsert=True - )) + operations.append(UpdateOne({"key": key, "feature_type": feature_type}, {"$set": data}, upsert=True)) if operations: try: self.collection.bulk_write(operations) except Exception as e: logger.error(f"Failed to bulk store embeddings in MongoDB: {e}") else: - for key, embedding, metadata in zip(keys, embeddings, metadata_list): + for key, embedding, metadata in zip(keys, embeddings, metadata_list, strict=False): self.store_embedding(key, embedding, feature_type, metadata) # Update cache for all - for key, embedding, metadata in zip(keys, embeddings, metadata_list): + for key, embedding, metadata in zip(keys, embeddings, metadata_list, strict=False): cache_key = f"{feature_type}:{key}" self._cache[cache_key] = {"embedding": embedding, "metadata": metadata or {}} @@ -195,9 +190,9 @@ def store_batch( def retrieve_batch( self, - keys: List[str], + keys: list[str], feature_type: str = "molecule", - ) -> Dict[str, np.ndarray]: + ) -> dict[str, np.ndarray]: """ Retrieve multiple embeddings. @@ -209,7 +204,7 @@ def retrieve_batch( Dictionary mapping keys to embeddings """ results = {} - + # Check cache first and identify missing keys missing_keys = [] for key in keys: @@ -218,7 +213,7 @@ def retrieve_batch( results[key] = self._cache[cache_key]["embedding"] else: missing_keys.append(key) - + if not missing_keys: return results @@ -247,7 +242,7 @@ def clear_cache(self) -> None: self._cache.clear() logger.info("Cleared feature store cache") - def get_statistics(self) -> Dict[str, Any]: + def get_statistics(self) -> dict[str, Any]: """Get statistics about stored features.""" stats = { "cache_size": len(self._cache), diff --git a/drug_discovery/data/normalizer.py b/drug_discovery/data/normalizer.py index a0036cb1ba..8a5940557b 100644 --- a/drug_discovery/data/normalizer.py +++ b/drug_discovery/data/normalizer.py @@ -5,7 +5,6 @@ """ import logging -from typing import List, Set, Optional, Tuple import pandas as pd @@ -13,7 +12,7 @@ try: # pragma: no cover - optional dependency from rdkit import Chem # type: ignore - from rdkit.Chem import Descriptors, AllChem # type: ignore + from rdkit.Chem import AllChem, Descriptors # type: ignore except Exception: # pragma: no cover - default path in slim environments Chem = None # type: ignore Descriptors = None # type: ignore @@ -35,9 +34,9 @@ def __init__(self, remove_salts: bool = True, remove_duplicates: bool = True): """ self.remove_salts = remove_salts self.remove_duplicates = remove_duplicates - self._seen_inchikeys: Set[str] = set() + self._seen_inchikeys: set[str] = set() - def canonicalize_smiles(self, smiles: str) -> Optional[str]: + def canonicalize_smiles(self, smiles: str) -> str | None: """ Convert SMILES to canonical form. @@ -74,7 +73,7 @@ def canonicalize_smiles(self, smiles: str) -> Optional[str]: logger.warning(f"Failed to canonicalize SMILES '{smiles}': {e}") return None - def compute_inchikey(self, smiles: str) -> Optional[str]: + def compute_inchikey(self, smiles: str) -> str | None: """ Compute InChIKey for molecular uniqueness. @@ -154,7 +153,7 @@ def normalize_dataframe( normalized_data = [] - for idx, row in df.iterrows(): + for _idx, row in df.iterrows(): smiles = row[smiles_column] if pd.isna(smiles): @@ -209,16 +208,13 @@ def normalize_dataframe( normalized_data.append(new_row) result_df = pd.DataFrame(normalized_data) - logger.info( - f"Normalized {len(df)} -> {len(result_df)} molecules " - f"(removed {len(df) - len(result_df)})" - ) + logger.info(f"Normalized {len(df)} -> {len(result_df)} molecules " f"(removed {len(df) - len(result_df)})") return result_df def merge_datasets( self, - datasets: List[pd.DataFrame], + datasets: list[pd.DataFrame], smiles_column: str = "smiles", ) -> pd.DataFrame: """ @@ -246,7 +242,7 @@ def apply_filters( df: pd.DataFrame, smiles_column: str = "smiles", lipinski_filter: bool = True, - molecular_weight_range: Optional[Tuple[float, float]] = None, + molecular_weight_range: tuple[float, float] | None = None, ) -> pd.DataFrame: """ Apply drug-likeness filters to dataset. @@ -280,13 +276,8 @@ def apply_filters( if molecular_weight_range: min_mw, max_mw = molecular_weight_range - filtered = filtered[ - (filtered["mol_weight"] >= min_mw) & (filtered["mol_weight"] <= max_mw) - ] + filtered = filtered[(filtered["mol_weight"] >= min_mw) & (filtered["mol_weight"] <= max_mw)] - logger.info( - f"Applied filters: {len(df)} -> {len(filtered)} molecules " - f"(removed {len(df) - len(filtered)})" - ) + logger.info(f"Applied filters: {len(df)} -> {len(filtered)} molecules " f"(removed {len(df) - len(filtered)})") return filtered diff --git a/drug_discovery/data/pipeline.py b/drug_discovery/data/pipeline.py index b09916d510..e9921da74e 100644 --- a/drug_discovery/data/pipeline.py +++ b/drug_discovery/data/pipeline.py @@ -11,96 +11,146 @@ """ from __future__ import annotations -import logging, re + +import logging from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional, Tuple +from typing import Any + import numpy as np logger = logging.getLogger(__name__) + def is_valid_smiles_fast(smiles): - if not smiles or len(smiles) > 2000: return False + if not smiles or len(smiles) > 2000: + return False return smiles.count("(") == smiles.count(")") and smiles.count("[") == smiles.count("]") + def validate_smiles(smiles): try: from rdkit import Chem + mol = Chem.MolFromSmiles(smiles) return (True, Chem.MolToSmiles(mol, isomericSmiles=True)) if mol else (False, None) except ImportError: return is_valid_smiles_fast(smiles), smiles + def validate_batch(smiles_list): valid, invalid, canonical = [], [], [] for s in smiles_list: ok, c = validate_smiles(s) - if ok: valid.append(s); canonical.append(c) - else: invalid.append(s) + if ok: + valid.append(s) + canonical.append(c) + else: + invalid.append(s) uc = list(set(canonical)) - return {"total": len(smiles_list), "valid": len(valid), "invalid": len(invalid), - "valid_ratio": len(valid)/max(len(smiles_list),1), "unique_canonical": len(uc), - "duplicate_count": len(canonical)-len(uc), "invalid_smiles": invalid[:20]} + return { + "total": len(smiles_list), + "valid": len(valid), + "invalid": len(invalid), + "valid_ratio": len(valid) / max(len(smiles_list), 1), + "unique_canonical": len(uc), + "duplicate_count": len(canonical) - len(uc), + "invalid_smiles": invalid[:20], + } + def compute_descriptors(smiles): try: - from rdkit import Chem; from rdkit.Chem import Descriptors, Crippen + from rdkit import Chem + from rdkit.Chem import Crippen, Descriptors + mol = Chem.MolFromSmiles(smiles) - if mol is None: return None - return {"mol_weight": Descriptors.MolWt(mol), "logp": Crippen.MolLogP(mol), - "hbd": Descriptors.NumHDonors(mol), "hba": Descriptors.NumHAcceptors(mol), - "tpsa": Descriptors.TPSA(mol), "rotatable_bonds": Descriptors.NumRotatableBonds(mol), - "aromatic_rings": Descriptors.NumAromaticRings(mol), - "heavy_atoms": Descriptors.HeavyAtomCount(mol), "qed": Descriptors.qed(mol)} - except ImportError: return None + if mol is None: + return None + return { + "mol_weight": Descriptors.MolWt(mol), + "logp": Crippen.MolLogP(mol), + "hbd": Descriptors.NumHDonors(mol), + "hba": Descriptors.NumHAcceptors(mol), + "tpsa": Descriptors.TPSA(mol), + "rotatable_bonds": Descriptors.NumRotatableBonds(mol), + "aromatic_rings": Descriptors.NumAromaticRings(mol), + "heavy_atoms": Descriptors.HeavyAtomCount(mol), + "qed": Descriptors.qed(mol), + } + except ImportError: + return None + def lipinski_filter(desc): - checks = {"mw_le_500": desc.get("mol_weight",0)<=500, "logp_le_5": desc.get("logp",0)<=5, - "hbd_le_5": desc.get("hbd",0)<=5, "hba_le_10": desc.get("hba",0)<=10} + checks = { + "mw_le_500": desc.get("mol_weight", 0) <= 500, + "logp_le_5": desc.get("logp", 0) <= 5, + "hbd_le_5": desc.get("hbd", 0) <= 5, + "hba_le_10": desc.get("hba", 0) <= 10, + } v = sum(1 for c in checks.values() if not c) - return {"passes": v<=1, "violations": v, "checks": checks} + return {"passes": v <= 1, "violations": v, "checks": checks} + def compute_morgan_fingerprint(smiles, radius=2, nbits=2048): try: - from rdkit import Chem; from rdkit.Chem import AllChem + from rdkit import Chem + from rdkit.Chem import AllChem + mol = Chem.MolFromSmiles(smiles) - if mol is None: return None + if mol is None: + return None return np.array(AllChem.GetMorganFingerprintAsBitVect(mol, radius, nBits=nbits), dtype=np.float32) - except ImportError: return None + except ImportError: + return None + def tanimoto_similarity(fp1, fp2): - inter = np.sum(fp1*fp2); union = np.sum(fp1)+np.sum(fp2)-inter + inter = np.sum(fp1 * fp2) + union = np.sum(fp1) + np.sum(fp2) - inter return float(inter / max(union, 1e-12)) + def smiles_to_graph(smiles): try: - from rdkit import Chem; from rdkit.Chem import AllChem + from rdkit import Chem + from rdkit.Chem import AllChem + mol = Chem.MolFromSmiles(smiles) - if mol is None: return None + if mol is None: + return None mol = Chem.AddHs(mol) try: - AllChem.EmbedMolecule(mol, AllChem.ETKDG()); AllChem.MMFFOptimizeMolecule(mol, maxIters=200) + AllChem.EmbedMolecule(mol, AllChem.ETKDG()) + AllChem.MMFFOptimizeMolecule(mol, maxIters=200) conf = mol.GetConformer() pos = np.array([conf.GetAtomPosition(i) for i in range(mol.GetNumAtoms())], dtype=np.float32) - except: pos = np.zeros((mol.GetNumAtoms(), 3), dtype=np.float32) + except: + pos = np.zeros((mol.GetNumAtoms(), 3), dtype=np.float32) z = np.array([a.GetAtomicNum() for a in mol.GetAtoms()], dtype=np.int64) src, dst = [], [] for b in mol.GetBonds(): - i, j = b.GetBeginAtomIdx(), b.GetEndAtomIdx(); src.extend([i,j]); dst.extend([j,i]) - ei = np.array([src, dst], dtype=np.int64) if src else np.zeros((2,0), dtype=np.int64) + i, j = b.GetBeginAtomIdx(), b.GetEndAtomIdx() + src.extend([i, j]) + dst.extend([j, i]) + ei = np.array([src, dst], dtype=np.int64) if src else np.zeros((2, 0), dtype=np.int64) return {"z": z, "pos": pos, "edge_index": ei, "num_atoms": len(z)} - except ImportError: return None + except ImportError: + return None + @dataclass class MolecularDataset: - smiles: List[str] = field(default_factory=list) - targets: Optional[np.ndarray] = None - descriptors: Optional[List[Dict]] = None - fingerprints: Optional[np.ndarray] = None - graphs: Optional[List[Dict]] = None - metadata: Dict[str, Any] = field(default_factory=dict) + smiles: list[str] = field(default_factory=list) + targets: np.ndarray | None = None + descriptors: list[dict] | None = None + fingerprints: np.ndarray | None = None + graphs: list[dict] | None = None + metadata: dict[str, Any] = field(default_factory=dict) @property - def size(self): return len(self.smiles) + def size(self): + return len(self.smiles) def featurize(self, methods=None): methods = methods or ["descriptors", "fingerprints"] @@ -119,6 +169,6 @@ def quality_report(self): if self.descriptors: vd = [d for d in self.descriptors if d] if vd: - report["lipinski_pass_rate"] = sum(1 for d in vd if lipinski_filter(d)["passes"])/len(vd) + report["lipinski_pass_rate"] = sum(1 for d in vd if lipinski_filter(d)["passes"]) / len(vd) report["mean_qed"] = np.mean([d["qed"] for d in vd if "qed" in d]) return report diff --git a/drug_discovery/data/versioning.py b/drug_discovery/data/versioning.py index 726b7113fb..98a8c049dd 100644 --- a/drug_discovery/data/versioning.py +++ b/drug_discovery/data/versioning.py @@ -5,12 +5,13 @@ track data evolution over time. """ -import logging -import json import hashlib -from pathlib import Path -from typing import Dict, List, Optional, Any +import json +import logging from datetime import datetime +from pathlib import Path +from typing import Any + import pandas as pd logger = logging.getLogger(__name__) @@ -35,7 +36,7 @@ def __init__(self, versions_dir: str = "./data/versions"): def _load_manifest(self) -> None: """Load the version manifest.""" if self.manifest_path.exists(): - with open(self.manifest_path, "r") as f: + with open(self.manifest_path) as f: self.manifest = json.load(f) else: self.manifest = {"versions": []} @@ -55,8 +56,8 @@ def create_version( self, dataset: pd.DataFrame, version_name: str, - description: Optional[str] = None, - metadata: Optional[Dict] = None, + description: str | None = None, + metadata: dict | None = None, ) -> str: """ Create a new dataset version. @@ -98,7 +99,7 @@ def create_version( logger.info(f"Created dataset version: {version_id} ({version_name})") return version_id - def load_version(self, version_id: str) -> Optional[pd.DataFrame]: + def load_version(self, version_id: str) -> pd.DataFrame | None: """ Load a specific dataset version. @@ -129,7 +130,7 @@ def load_version(self, version_id: str) -> Optional[pd.DataFrame]: logger.info(f"Loaded dataset version: {version_id}") return dataset - def list_versions(self) -> List[Dict[str, Any]]: + def list_versions(self) -> list[dict[str, Any]]: """ List all available versions. @@ -138,7 +139,7 @@ def list_versions(self) -> List[Dict[str, Any]]: """ return self.manifest["versions"] - def get_latest_version(self) -> Optional[str]: + def get_latest_version(self) -> str | None: """ Get the ID of the most recent version. @@ -153,7 +154,7 @@ def compare_versions( self, version_id_1: str, version_id_2: str, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """ Compare two dataset versions. @@ -215,7 +216,7 @@ def tag_version(self, version_id: str, tag: str) -> bool: logger.error(f"Version not found: {version_id}") return False - def get_versions_by_tag(self, tag: str) -> List[str]: + def get_versions_by_tag(self, tag: str) -> list[str]: """ Get all version IDs with a specific tag. diff --git a/drug_discovery/diffusion/__init__.py b/drug_discovery/diffusion/__init__.py index 5a79ee9da3..8ed193d0c5 100644 --- a/drug_discovery/diffusion/__init__.py +++ b/drug_discovery/diffusion/__init__.py @@ -22,10 +22,10 @@ ) __all__ = [ - "EquivariantDiffusionModel", "DiffusionConfig", "DiffusionResult", + "EquivariantDiffusionModel", + "GeneratedMolecule", "PocketAwareGenerator", "PocketContext", - "GeneratedMolecule", ] diff --git a/drug_discovery/diffusion/diffusion_model.py b/drug_discovery/diffusion/diffusion_model.py index 8f1c3b296c..fb3db05795 100644 --- a/drug_discovery/diffusion/diffusion_model.py +++ b/drug_discovery/diffusion/diffusion_model.py @@ -409,7 +409,7 @@ def generate( generated_coords=generated_coords, binding_scores=binding_scores, sa_scores=sa_scores, - quality_scores=[sa - abs(b - 0.5) for sa, b in zip(sa_scores, binding_scores)], + quality_scores=[sa - abs(b - 0.5) for sa, b in zip(sa_scores, binding_scores, strict=False)], n_valid=len(valid_smiles), n_unique=len(unique_smiles), success=True, diff --git a/drug_discovery/diffusion/pocket_generator.py b/drug_discovery/diffusion/pocket_generator.py index 6ca4d66bb0..d7f89777cb 100644 --- a/drug_discovery/diffusion/pocket_generator.py +++ b/drug_discovery/diffusion/pocket_generator.py @@ -172,10 +172,7 @@ def generate( # Extract pocket information pocket_coords = pocket_context.pocket_coords - if pocket_coords is not None: - center = pocket_coords.mean(axis=0) - else: - center = np.zeros(3) + center = pocket_coords.mean(axis=0) if pocket_coords is not None else np.zeros(3) # Generate candidates candidates = [] @@ -183,7 +180,7 @@ def generate( # Drug-like scaffolds for generation scaffolds = self._get_scaffolds_for_pocket(pocket_context) - for i in range(n_molecules): + for _i in range(n_molecules): # Select scaffold scaffold = scaffolds[np.random.randint(0, len(scaffolds))] diff --git a/drug_discovery/docking/__init__.py b/drug_discovery/docking/__init__.py index 658c6c1d2a..8014b957b1 100644 --- a/drug_discovery/docking/__init__.py +++ b/drug_discovery/docking/__init__.py @@ -1,2 +1,2 @@ +from .diffdock_placeholder import DiffDockStub from .vina_wrapper import VinaDocker -from .diffdock_placeholder import DiffDockStub \ No newline at end of file diff --git a/drug_discovery/docking/diffdock_placeholder.py b/drug_discovery/docking/diffdock_placeholder.py index 040aa2d054..dcc7093387 100644 --- a/drug_discovery/docking/diffdock_placeholder.py +++ b/drug_discovery/docking/diffdock_placeholder.py @@ -1,9 +1,4 @@ def run_diffdock(ligand_smiles: str, protein_pdb: str, num_poses: int = 10) -> dict: # Stub for external/DiffDock integration # In full: subprocess.call(['python', 'external/DiffDock/run.py', ...]) - return { - 'poses': [], - 'affinity': -7.5, # kcal/mol stub - 'rmsd': 1.2, - 'success': True - } \ No newline at end of file + return {"poses": [], "affinity": -7.5, "rmsd": 1.2, "success": True} # kcal/mol stub diff --git a/drug_discovery/docking/vina_wrapper.py b/drug_discovery/docking/vina_wrapper.py index 629becb73d..a5a8b5cb9c 100644 --- a/drug_discovery/docking/vina_wrapper.py +++ b/drug_discovery/docking/vina_wrapper.py @@ -1,25 +1,25 @@ -from vina import Vina -from rdkit import Chem -from rdkit.Chem import AllChem import os +from vina import Vina + + class VinaDocker: def __init__(self, exhaustiveness: int = 8): - self.vina = Vina(sf_name='vina') + self.vina = Vina(sf_name="vina") self.exhaustiveness = exhaustiveness def prepare_receptor(self, pdb_path: str) -> str: # Simple prep: assume clean PDB - receptor_pdbqt = pdb_path.replace('.pdb', '_receptor.pdbqt') - os.system(f'obabel {pdb_path} -O {receptor_pdbqt} -xh') # Add H + receptor_pdbqt = pdb_path.replace(".pdb", "_receptor.pdbqt") + os.system(f"obabel {pdb_path} -O {receptor_pdbqt} -xh") # Add H return receptor_pdbqt def prepare_ligands(self, sdf_path: str) -> str: - ligands_pdbqt = sdf_path.replace('.sdf', '_ligands.pdbqt') - os.system(f'obabel {sdf_path} -O {ligands_pdbqt} -xh') + ligands_pdbqt = sdf_path.replace(".sdf", "_ligands.pdbqt") + os.system(f"obabel {sdf_path} -O {ligands_pdbqt} -xh") return ligands_pdbqt - def dock(self, receptor_pdb: str, ligands_sdf: str, center: tuple = (0,0,0), size: tuple = (20,20,20)) -> dict: + def dock(self, receptor_pdb: str, ligands_sdf: str, center: tuple = (0, 0, 0), size: tuple = (20, 20, 20)) -> dict: receptor = self.prepare_receptor(receptor_pdb) ligands = self.prepare_ligands(ligands_sdf) self.vina.set_receptor(receptor) @@ -27,4 +27,4 @@ def dock(self, receptor_pdb: str, ligands_sdf: str, center: tuple = (0,0,0), siz self.vina.compute_vina_maps(center=center, box_size=list(size)) self.vina.dock(exhaustiveness=self.exhaustiveness, n_poses=9) affinities = self.vina.energies() - return {'affinities': affinities, 'poses': self.vina.poses()} \ No newline at end of file + return {"affinities": affinities, "poses": self.vina.poses()} diff --git a/drug_discovery/drug_repurposing/similarity_search.py b/drug_discovery/drug_repurposing/similarity_search.py index 0fe865aefd..959f5550ec 100644 --- a/drug_discovery/drug_repurposing/similarity_search.py +++ b/drug_discovery/drug_repurposing/similarity_search.py @@ -1,34 +1,31 @@ -import numpy as np +from typing import Any + from rdkit import Chem, DataStructs from rdkit.Chem import AllChem -from typing import List, Tuple, Dict + class DrugSimilaritySearch: """ Find drugs similar to a query molecule for potential repurposing. """ - - def __init__(self, drug_library_smiles: List[str], drug_names: List[str] = None): + + def __init__(self, drug_library_smiles: list[str], drug_names: list[str] | None = None): self.library_smiles = drug_library_smiles self.library_mols = [Chem.MolFromSmiles(s) for s in drug_library_smiles] self.library_fps = [AllChem.GetMorganFingerprintAsBitVect(m, 2, nBits=2048) for m in self.library_mols if m] self.drug_names = drug_names or [f"Drug_{i}" for i in range(len(drug_library_smiles))] - def find_similar(self, query_smiles: str, threshold: float = 0.7) -> List[Dict[str, Any]]: + def find_similar(self, query_smiles: str, threshold: float = 0.7) -> list[dict[str, Any]]: query_mol = Chem.MolFromSmiles(query_smiles) if not query_mol: return [] - + query_fp = AllChem.GetMorganFingerprintAsBitVect(query_mol, 2, nBits=2048) - + results = [] for i, fp in enumerate(self.library_fps): similarity = DataStructs.TanimotoSimilarity(query_fp, fp) if similarity >= threshold: - results.append({ - "name": self.drug_names[i], - "smiles": self.library_smiles[i], - "similarity": similarity - }) - + results.append({"name": self.drug_names[i], "smiles": self.library_smiles[i], "similarity": similarity}) + return sorted(results, key=lambda x: x["similarity"], reverse=True) diff --git a/drug_discovery/drug_repurposing/target_prediction.py b/drug_discovery/drug_repurposing/target_prediction.py index df94637cf3..10f01ea721 100644 --- a/drug_discovery/drug_repurposing/target_prediction.py +++ b/drug_discovery/drug_repurposing/target_prediction.py @@ -1,30 +1,27 @@ -from typing import List, Dict, Any -import torch +from typing import Any + class ReverseScreening: """ Predict potential targets for a given drug (reverse screening). """ - - def __init__(self, target_models: Dict[str, Any]): + + def __init__(self, target_models: dict[str, Any]): """ Args: target_models: Dictionary mapping target names to their activity prediction models. """ self.target_models = target_models - - def predict_targets(self, smiles: str) -> List[Dict[str, Any]]: + + def predict_targets(self, smiles: str) -> list[dict[str, Any]]: """ Predict activity against all registered targets. """ results = [] - for target_name, model in self.target_models.items(): + for target_name, _model in self.target_models.items(): # In a real scenario, we'd use the model to predict activity # score = model.predict(smiles) - score = 0.5 # Placeholder - results.append({ - "target": target_name, - "score": score - }) - + score = 0.5 # Placeholder + results.append({"target": target_name, "score": score}) + return sorted(results, key=lambda x: x["score"], reverse=True) diff --git a/drug_discovery/drugmaking/__init__.py b/drug_discovery/drugmaking/__init__.py index 98b1d2b939..04d850f841 100644 --- a/drug_discovery/drugmaking/__init__.py +++ b/drug_discovery/drugmaking/__init__.py @@ -34,15 +34,15 @@ from .vae_generator import DeliveryVAE as DeliveryVAE __all__ = [ - "CustomDrugmakingModule", - "CompoundTestResult", + "LNP", "CandidateResult", - "OptimizationConfig", + "CompoundTestResult", "CounterSubstanceFinder", "CounterSubstanceResult", - "LNP", - "PolymericSystem", + "CustomDrugmakingModule", + "DeliveryGenerator", "DeliverySystem", "DeliveryVAE", - "DeliveryGenerator", + "OptimizationConfig", + "PolymericSystem", ] diff --git a/drug_discovery/drugmaking/process.py b/drug_discovery/drugmaking/process.py index edcbae42d9..445de89902 100644 --- a/drug_discovery/drugmaking/process.py +++ b/drug_discovery/drugmaking/process.py @@ -255,7 +255,7 @@ def compute_composite_score(self, config: OptimizationConfig) -> float: # Weighted effectiveness effectiveness = 0.0 - for i, name in enumerate(config.objective_names): + for _i, name in enumerate(config.objective_names): if name in ["potency", "selectivity", "solubility"]: val = self.objectives.get(name, 0.5) effectiveness += val * config.effectiveness_weight diff --git a/drug_discovery/drugmaking/risk_mitigation.py b/drug_discovery/drugmaking/risk_mitigation.py index 8fdafbe46b..1d56ccb206 100644 --- a/drug_discovery/drugmaking/risk_mitigation.py +++ b/drug_discovery/drugmaking/risk_mitigation.py @@ -363,9 +363,8 @@ def _compute_mechanism_score(self, smiles: str, target_toxicity: str | None) -> return NEUTRALIZATION_RULES[target_toxicity]["chemistry"] if "carboxylic_acid" in groups or "sulfonate" in groups: return "acid-base neutralization" - if "amine" in groups: - if target_toxicity in ["acid_toxicity", "metal_toxicity"]: - return "acid-base neutralization" if "acid" in target_toxicity else "chelation" + if "amine" in groups and target_toxicity in ["acid_toxicity", "metal_toxicity"]: + return "acid-base neutralization" if "acid" in target_toxicity else "chelation" if len(groups) >= 2 and "amine" in groups and "carboxylic_acid" in groups: return "chelation" return "unknown" diff --git a/drug_discovery/evaluation/__init__.py b/drug_discovery/evaluation/__init__.py index df396ff91d..3ea61307cb 100644 --- a/drug_discovery/evaluation/__init__.py +++ b/drug_discovery/evaluation/__init__.py @@ -27,9 +27,9 @@ __all__.extend( [ - "MCDropoutPredictor", - "DeepEnsemble", "ConformalPredictor", + "DeepEnsemble", + "MCDropoutPredictor", "UncertaintyConfig", "expected_calibration_error", "regression_calibration_error", @@ -52,7 +52,7 @@ compute_admet_profile as compute_admet_profile, ) - __all__.extend(["AdvancedADMETPredictor", "ADMETConfig", "ADMET_ENDPOINTS", "compute_admet_profile"]) + __all__.extend(["ADMET_ENDPOINTS", "ADMETConfig", "AdvancedADMETPredictor", "compute_admet_profile"]) except ImportError as e: logger.debug(f"Advanced ADMET not available: {e}") @@ -91,6 +91,6 @@ from drug_discovery.evaluation.herg_predictor import HERGPredictor as HERGPredictor from drug_discovery.evaluation.herg_predictor import predict_herg as predict_herg - __all__.extend(["HERGPredictor", "HERGPrediction", "predict_herg"]) + __all__.extend(["HERGPrediction", "HERGPredictor", "predict_herg"]) except ImportError as e: logger.debug(f"hERG predictor not available: {e}") diff --git a/drug_discovery/evaluation/advanced_admet.py b/drug_discovery/evaluation/advanced_admet.py index 764fd03bde..85f97d35b0 100644 --- a/drug_discovery/evaluation/advanced_admet.py +++ b/drug_discovery/evaluation/advanced_admet.py @@ -15,13 +15,12 @@ from __future__ import annotations import logging -import os from dataclasses import dataclass, field +import numpy as np import torch import torch.nn as nn import torch.nn.functional as F -import numpy as np logger = logging.getLogger(__name__) @@ -197,7 +196,7 @@ def predict_batch(self, batch_data: dict) -> dict[str, torch.Tensor]: edge_index=batch_data["edge_index"], smiles_tokens=batch_data["smiles_tokens"], batch=batch_data.get("batch"), - smiles_mask=batch_data.get("smiles_mask") + smiles_mask=batch_data.get("smiles_mask"), ) @@ -224,7 +223,10 @@ def compute_admet_profile(predictions: dict[str, torch.Tensor]) -> dict[str, dic } else: val = pred.cpu().numpy() - profile[name] = {"value": val.tolist() if val.ndim > 0 else round(val.item(), 4), "unit": info.get("unit", "")} + profile[name] = { + "value": val.tolist() if val.ndim > 0 else round(val.item(), 4), + "unit": info.get("unit", ""), + } return profile @@ -234,9 +236,11 @@ def compute_admet_profile(predictions: dict[str, torch.Tensor]) -> dict[str, dic ray = None if ray: + @ray.remote(num_gpus=1 if torch.cuda.is_available() else 0) class RayADMETPredictor: """Ray-wrapped ADMET predictor for distributed inference.""" + def __init__(self, config: ADMETConfig): self.model = AdvancedADMETPredictor(config) self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -249,7 +253,9 @@ def predict(self, batch_data: dict) -> dict[str, np.ndarray]: with torch.no_grad(): preds = self.model.predict_batch(input_data) return {k: v.cpu().numpy() for k, v in preds.items()} + else: + class RayADMETPredictor: def __init__(self, *args, **kwargs): raise ImportError("Ray not installed.") diff --git a/drug_discovery/evaluation/herg_predictor.py b/drug_discovery/evaluation/herg_predictor.py index a0fa98da7e..b6d1500dbd 100644 --- a/drug_discovery/evaluation/herg_predictor.py +++ b/drug_discovery/evaluation/herg_predictor.py @@ -44,14 +44,14 @@ class HERGPrediction: class HERGPredictor: """QSAR model for hERG potassium channel inhibition. - + Based on literature QSAR models (e.g., Recanatini et al., literature review): - Lipophilicity (logP): positive correlation with hERG inhibition - Molecular weight: weak positive correlation - TPSA: negative correlation (more polar = less hERG active) - Basic nitrogen count: positive correlation (mimics piperidine in known hERG blockers) - H-bond donors: mixed effect - + The model is parametrized to allow fine-tuning without hardcoding decision boundaries. """ @@ -59,27 +59,24 @@ def __init__( self, # QSAR coefficients (learned from data, tunable parameters) logp_coeff: float = 0.40, # logP contribution to hERG inhibition (reduced for more specificity) - mw_coeff: float = 0.08, # MW contribution (weak) + mw_coeff: float = 0.08, # MW contribution (weak) tpsa_coeff: float = -0.30, # TPSA contribution (negative = less polar is more hERG active) basic_n_coeff: float = 0.20, # Basic nitrogen count contribution - hbd_coeff: float = -0.05, # H-bond donors (weak/negative) + hbd_coeff: float = -0.05, # H-bond donors (weak/negative) aromatic_rings_coeff: float = 0.15, # Aromatic rings (planar = hERG risk) - # IC50 estimation parameters ic50_baseline_nM: float = 5000.0, # Baseline IC50 for reference compound ic50_potency_range: tuple[float, float] = (100.0, 50000.0), # [most potent, least potent] - # Risk classification thresholds (CiPA-oriented) cipa_low_threshold: float = 0.25, # P(inhibitor) threshold for low risk (increased for specificity) cipa_cat2_threshold: float = 0.50, # Threshold for category 2 (intermediate) cipa_high_threshold: float = 0.75, # Threshold for high risk - # QTc prolongation parameters qtc_threshold_ms: float = 60.0, # Clinical QTc change threshold ic50_for_qtc_risk_nM: float = 10000.0, # IC50 where QTc risk becomes moderate ): """Initialize hERG predictor with parametrized QSAR coefficients. - + Args: logp_coeff: Lipophilicity coefficient (positive = higher logP increases hERG risk) mw_coeff: Molecular weight coefficient @@ -102,69 +99,69 @@ def __init__( self.basic_n_coeff = basic_n_coeff self.hbd_coeff = hbd_coeff self.aromatic_rings_coeff = aromatic_rings_coeff - + # IC50 parameters self.ic50_baseline_nM = ic50_baseline_nM self.ic50_potency_min, self.ic50_potency_max = ic50_potency_range - + # Risk classification parameters self.cipa_low_threshold = cipa_low_threshold self.cipa_cat2_threshold = cipa_cat2_threshold self.cipa_high_threshold = cipa_high_threshold - + # QTc parameters self.qtc_threshold_ms = qtc_threshold_ms self.ic50_for_qtc_risk_nM = ic50_for_qtc_risk_nM def predict(self, smiles: str, calibrate: bool = True) -> HERGPrediction: """Predict hERG inhibition probability for given SMILES string. - + Args: smiles: SMILES representation of molecule calibrate: Whether to apply probability calibration - + Returns: HERGPrediction with inhibition probability, IC50 estimate, and risk classification """ # Calculate molecular properties props = self._calculate_properties(smiles) - + # Calculate raw QSAR score (logit scale, not bounded) qsar_score = self._calculate_qsar_score(props) - + # Convert to probability via logistic function inhibition_prob = self._logistic(qsar_score) - + # Apply calibration if requested if calibrate: inhibition_prob = self._calibrate_probability(inhibition_prob) - + # Bound to [0, 1] inhibition_prob = max(0.0, min(1.0, inhibition_prob)) - + # Estimate IC50 from probability ic50_estimate = self._estimate_ic50(inhibition_prob, props) - + # Calculate IC50 confidence interval ic50_low, ic50_high = self._ic50_confidence_interval(ic50_estimate, inhibition_prob) - + # Classify CiPA risk category cipa_risk = self._classify_cipa_risk(inhibition_prob) - + # Classify QTc prolongation risk qtc_risk = self._classify_qtc_risk(inhibition_prob, ic50_estimate) - + # Probability confidence interval prob_std = self._estimate_probability_std(inhibition_prob) prob_low = max(0.0, inhibition_prob - 1.96 * prob_std) prob_high = min(1.0, inhibition_prob + 1.96 * prob_std) - + # Model confidence based on property uncertainty model_confidence = self._calculate_model_confidence(props) - + # Identify key concerns concerns = self._identify_concerns(props, inhibition_prob, ic50_estimate) - + return HERGPrediction( smiles=smiles, inhibition_probability=float(inhibition_prob), @@ -194,7 +191,7 @@ def _calculate_properties(self, smiles: str) -> dict[str, float]: "rotatable_bonds": int(Descriptors.NumRotatableBonds(mol)), "hba": int(rdMolDescriptors.CalcNumHBA(mol)), } - + # Fallback: estimate from SMILES string return self._estimate_properties_from_smiles(smiles) @@ -223,7 +220,7 @@ def _estimate_properties_from_smiles(self, smiles: str) -> dict[str, float]: # Heuristic estimation upper_count = sum(1 for c in smiles if c.isupper() and c != "N") n_count = smiles.count("N") - + return { "logp": 1.0 + upper_count * 0.3 + n_count * (-0.2), # Higher C count = higher logP, N = lower "mw": upper_count * 14.0 + n_count * 14.0 + 100.0, @@ -243,7 +240,7 @@ def _calculate_qsar_score(self, props: dict[str, float]) -> float: basic_n = props.get("basic_n_count", 0) hbd = props.get("hbd", 1) aromatic = props.get("aromatic_rings", 0) - + # Normalize features to similar scales logp_term = self.logp_coeff * (logp - 2.0) # Center at typical logP of 2 mw_term = self.mw_coeff * (mw - 250.0) / 100.0 # Normalize by typical MW ~250-400 @@ -251,10 +248,10 @@ def _calculate_qsar_score(self, props: dict[str, float]) -> float: basic_n_term = self.basic_n_coeff * basic_n hbd_term = self.hbd_coeff * hbd aromatic_term = self.aromatic_rings_coeff * aromatic - + # Sum components with baseline offset (slightly negative to prefer low risk) score = -0.5 + logp_term + mw_term + tpsa_term + basic_n_term + hbd_term + aromatic_term - + return score def _logistic(self, x: float) -> float: @@ -266,7 +263,7 @@ def _logistic(self, x: float) -> float: def _calibrate_probability(self, prob: float) -> float: """Apply calibration to predicted probability. - + Calibration corrects for training data bias if the model was trained on a different population than deployment data. """ @@ -277,23 +274,23 @@ def _calibrate_probability(self, prob: float) -> float: def _estimate_ic50(self, inhibition_prob: float, props: dict[str, float]) -> float: """Estimate IC50 (in nanoMolar) from inhibition probability and properties. - + Uses inverse relationship: higher probability = more potent (lower IC50). """ if inhibition_prob < 0.05: # Not an inhibitor -> very high IC50 (inactive) return self.ic50_potency_max - + # Log-linear relationship between probability and IC50 # prob = 1 / (1 + (IC50 / IC50_ref)^n) # Rearranging: IC50 = IC50_ref * (1/prob - 1)^(1/n) slope = 1.5 # Hill coefficient - + if inhibition_prob > 0.95: return self.ic50_potency_min - + ic50 = self.ic50_baseline_nM * ((1.0 / inhibition_prob - 1.0) ** (1.0 / slope)) - + # Bound to plausible range return max(self.ic50_potency_min, min(self.ic50_potency_max, ic50)) @@ -301,19 +298,19 @@ def _ic50_confidence_interval(self, ic50: float, prob: float) -> tuple[float, fl """Calculate 90% confidence interval for IC50 estimate.""" # Uncertainty increases near 50% probability uncertainty_factor = 2.0 * prob * (1.0 - prob) - + # Log-scale confidence interval log_ic50 = math.log10(ic50) log_std = 0.3 * uncertainty_factor # Standard deviation in log scale - + log_low = log_ic50 - 1.645 * log_std log_high = log_ic50 + 1.645 * log_std - - return (10 ** log_low, 10 ** log_high) + + return (10**log_low, 10**log_high) def _classify_cipa_risk(self, prob: float) -> str: """Classify CiPA (Comprehensive in vitro Proarrhythmia Assay) risk category. - + - Low: p < 0.20 - Category 2: 0.20 <= p < 0.70 - Category 3 (High): p >= 0.70 @@ -327,17 +324,17 @@ def _classify_cipa_risk(self, prob: float) -> str: def _classify_qtc_risk(self, prob: float, ic50: float) -> str: """Classify QTc prolongation risk based on hERG inhibition. - + Risk depends on both: 1. Probability of hERG inhibition 2. Potency (IC50) - more potent = more risk - + Clinical significance: QTc change > 30 ms is concerning. """ # Combined risk score potency_risk = max(0.0, 1.0 - (ic50 / self.ic50_for_qtc_risk_nM)) combined_risk = 0.6 * prob + 0.4 * potency_risk - + if combined_risk < 0.15: return "very_low" elif combined_risk < 0.35: @@ -349,7 +346,7 @@ def _classify_qtc_risk(self, prob: float, ic50: float) -> str: def _estimate_probability_std(self, prob: float) -> float: """Estimate standard deviation of probability estimate. - + Uncertainty is highest near 50% (maximum entropy). """ # Binomial-like variance: p(1-p) @@ -361,51 +358,49 @@ def _calculate_model_confidence(self, props: dict[str, float]) -> float: # Higher confidence for properties in typical ranges mw = props.get("mw", 300.0) logp = props.get("logp", 2.0) - + # Penalty for extrapolation beyond typical drug ranges mw_penalty = 0.0 if 150 <= mw <= 600 else 0.2 logp_penalty = 0.0 if -1 <= logp <= 6 else 0.2 - + base_confidence = 0.8 return max(0.5, base_confidence - mw_penalty - logp_penalty) - def _identify_concerns( - self, props: dict[str, float], prob: float, ic50: float - ) -> list[str]: + def _identify_concerns(self, props: dict[str, float], prob: float, ic50: float) -> list[str]: """Identify key structural/property concerns contributing to hERG risk.""" concerns = [] - + logp = props.get("logp", 2.0) if logp > 4.0: concerns.append(f"Very high lipophilicity (logP={logp:.1f})") elif logp > 3.0: concerns.append(f"High lipophilicity (logP={logp:.1f})") - + tpsa = props.get("tpsa", 60.0) if tpsa < 40: concerns.append(f"Low TPSA indicates poor aqueous solubility ({tpsa:.0f})") - + basic_n = props.get("basic_n_count", 0) if basic_n >= 2: concerns.append(f"Multiple basic nitrogens ({basic_n}) can increase hERG binding") - + if prob > 0.5: concerns.append(f"High hERG inhibition probability ({prob:.1%})") - + if ic50 < 1000: concerns.append(f"Potent hERG inhibition (IC50={ic50:.0f} nM < 1 µM)") - + return concerns # Convenience functions def predict_herg(smiles: str, use_defaults: bool = True) -> HERGPrediction: """Quick hERG prediction using default model parameters. - + Args: smiles: SMILES string of molecule use_defaults: If True, use default QSAR coefficients; if False, requires tuning - + Returns: HERGPrediction object with full inhibition assessment """ diff --git a/drug_discovery/evaluation/predictor.py b/drug_discovery/evaluation/predictor.py index fbf37a78fe..e7b2c7acbe 100644 --- a/drug_discovery/evaluation/predictor.py +++ b/drug_discovery/evaluation/predictor.py @@ -345,10 +345,7 @@ def expected_calibration_error_regression( # Normalize uncertainty to [0, 1] for comparability across scales. unc_min = float(np.min(unc_v)) unc_max = float(np.max(unc_v)) - if unc_max - unc_min > 1e-12: - unc_norm = (unc_v - unc_min) / (unc_max - unc_min) - else: - unc_norm = np.zeros_like(unc_v) + unc_norm = (unc_v - unc_min) / (unc_max - unc_min) if unc_max - unc_min > 1e-12 else np.zeros_like(unc_v) abs_err = np.abs(true_v - pred_v) err_max = float(np.max(abs_err)) diff --git a/drug_discovery/evaluation/swissadme_proxy.py b/drug_discovery/evaluation/swissadme_proxy.py index 88797c74ad..61d788f1b8 100644 --- a/drug_discovery/evaluation/swissadme_proxy.py +++ b/drug_discovery/evaluation/swissadme_proxy.py @@ -1,7 +1,8 @@ from __future__ import annotations -from dataclasses import dataclass, field import re +from dataclasses import dataclass, field + @dataclass class SwissADMEResult: @@ -34,8 +35,12 @@ def _estimate_from_smiles(smiles: str) -> dict[str, float | int | bool]: triple_bonds = smiles.count("#") halogens = sum(smiles.count(atom) for atom in ("F", "Cl", "Br", "I")) - molecular_weight = float(max(0.0, 12.0 * carbon_like + 14.0 * hetero_atoms + 19.0 * halogens + 1.0 * smiles.count("H"))) - logp = round(0.54 * carbon_like - 1.35 * hetero_polar - 0.08 * ring_closures + 0.12 * double_bonds + 0.16 * triple_bonds, 2) + molecular_weight = float( + max(0.0, 12.0 * carbon_like + 14.0 * hetero_atoms + 19.0 * halogens + 1.0 * smiles.count("H")) + ) + logp = round( + 0.54 * carbon_like - 1.35 * hetero_polar - 0.08 * ring_closures + 0.12 * double_bonds + 0.16 * triple_bonds, 2 + ) hbd = int(min(10, hetero_polar)) hba = int(min(15, hetero_polar + halogens)) tpsa = round(11.0 * hetero_polar + 2.0 * ring_closures + 1.5 * double_bonds, 1) @@ -108,9 +113,22 @@ def _rdkit_profile(self, smiles: str) -> dict[str, float | int | bool] | None: "tpsa": tpsa, "rotatable_bonds": rotatable_bonds, "ring_closures": ring_closures, - "gi_absorption": round(max(0.0, min(1.0, 0.92 - 0.35 * _safe_div(tpsa, 140.0) - 0.06 * max(0.0, logp - 3.0))), 3), + "gi_absorption": round( + max(0.0, min(1.0, 0.92 - 0.35 * _safe_div(tpsa, 140.0) - 0.06 * max(0.0, logp - 3.0))), 3 + ), "bbb_permeant": molecular_weight <= 450 and tpsa <= 90 and logp <= 4.5, - "cyp_inhib": round(max(0.0, min(1.0, 0.15 + 0.12 * max(0.0, logp) + 0.05 * max(0, sum(1 for atom in mol.GetAtoms() if atom.GetSymbol() in {"Cl", "Br", "I", "F"}) - 1)))), + "cyp_inhib": round( + max( + 0.0, + min( + 1.0, + 0.15 + + 0.12 * max(0.0, logp) + + 0.05 + * max(0, sum(1 for atom in mol.GetAtoms() if atom.GetSymbol() in {"Cl", "Br", "I", "F"}) - 1), + ), + ) + ), "lipinski_ok": lipinski_violations <= 1, "veber_ok": rotatable_bonds <= 10 and tpsa <= 140, } @@ -152,4 +170,4 @@ def predict(self, smiles: str) -> SwissADMEResult: developable=developable, rule_checks=rule_checks, notes=notes, - ) \ No newline at end of file + ) diff --git a/drug_discovery/evaluation/uncertainty.py b/drug_discovery/evaluation/uncertainty.py index c98ecc459b..98632fed5f 100644 --- a/drug_discovery/evaluation/uncertainty.py +++ b/drug_discovery/evaluation/uncertainty.py @@ -140,8 +140,8 @@ def regression_calibration_error(preds, true_vals, uncertainties, bins=10): for lev in levels: z = 1.96 * lev # Simplified z-score coverages.append(float((errors <= z * uncertainties).mean())) - cal_errs = [abs(o - e) for o, e in zip(coverages, levels)] + cal_errs = [abs(o - e) for o, e in zip(coverages, levels, strict=False)] return { "mean_calibration_error": float(np.mean(cal_errs)), - "interval_coverages": dict(zip([f"{level_val:.1f}" for level_val in levels], coverages)), + "interval_coverages": dict(zip([f"{level_val:.1f}" for level_val in levels], coverages, strict=False)), } diff --git a/drug_discovery/explainability/__init__.py b/drug_discovery/explainability/__init__.py index f83544ff3a..392cdcc917 100644 --- a/drug_discovery/explainability/__init__.py +++ b/drug_discovery/explainability/__init__.py @@ -3,7 +3,7 @@ Methods for interpreting model predictions (XAI). """ -from .graph_explainer import GraphExplainer from .fingerprint_explainer import FingerprintExplainer +from .graph_explainer import GraphExplainer -__all__ = ["GraphExplainer", "FingerprintExplainer"] +__all__ = ["FingerprintExplainer", "GraphExplainer"] diff --git a/drug_discovery/explainability/fingerprint_explainer.py b/drug_discovery/explainability/fingerprint_explainer.py index 03b220b358..cda711ff6e 100644 --- a/drug_discovery/explainability/fingerprint_explainer.py +++ b/drug_discovery/explainability/fingerprint_explainer.py @@ -1,13 +1,15 @@ +from typing import Any + +import numpy as np import torch import torch.nn as nn -import numpy as np -from typing import Any, Dict, List, Optional, Tuple + class FingerprintExplainer: """ Explainer for fingerprint-based models. """ - + def __init__(self, model: nn.Module, device: str = "cpu"): self.model = model.to(device) self.model.eval() @@ -19,27 +21,27 @@ def attribute_bits(self, fingerprint: torch.Tensor) -> torch.Tensor: """ fingerprint = fingerprint.to(self.device).float() fingerprint.requires_grad = True - + output = self.model(fingerprint) score = output.sum() score.backward() - + attribution = fingerprint.grad return attribution - def explain_prediction(self, fingerprint: np.ndarray) -> Dict[str, Any]: + def explain_prediction(self, fingerprint: np.ndarray) -> dict[str, Any]: """ Explain a prediction by identifying key bits in the fingerprint. """ fp_tensor = torch.from_numpy(fingerprint) attributions = self.attribute_bits(fp_tensor) - + # Identify most influential bits attr_np = attributions.detach().cpu().numpy().flatten() top_bits = np.argsort(np.abs(attr_np))[::-1][:10] - + return { "bit_attributions": attr_np.tolist(), "top_influential_bits": top_bits.tolist(), - "explanation_method": "gradient" + "explanation_method": "gradient", } diff --git a/drug_discovery/explainability/graph_explainer.py b/drug_discovery/explainability/graph_explainer.py index 0b729a5561..73af2b5143 100644 --- a/drug_discovery/explainability/graph_explainer.py +++ b/drug_discovery/explainability/graph_explainer.py @@ -1,13 +1,15 @@ +from typing import Any + import torch import torch.nn as nn -from typing import Any, Dict, List, Optional, Tuple from torch_geometric.data import Data + class GraphExplainer: """ Explainer for Graph Neural Network models. """ - + def __init__(self, model: nn.Module, device: str = "cpu"): self.model = model.to(device) self.model.eval() @@ -20,31 +22,28 @@ def attribute_nodes(self, data: Data, target_class: int = 0) -> torch.Tensor: """ data = data.to(self.device) data.x.requires_grad = True - + output = self.model(data) - - if output.dim() > 1 and output.size(1) > 1: - score = output[0, target_class] - else: - score = output[0] - + + score = output[0, target_class] if output.dim() > 1 and output.size(1) > 1 else output[0] + score.backward() - + # Attribution is the norm of the gradient attribution = data.x.grad.norm(dim=1) return attribution - def explain_prediction(self, data: Data) -> Dict[str, Any]: + def explain_prediction(self, data: Data) -> dict[str, Any]: """ Generate a full explanation for a graph prediction. """ attributions = self.attribute_nodes(data) - + # Get top-k important nodes top_indices = torch.argsort(attributions, descending=True)[:5] - + return { "node_attributions": attributions.tolist(), "top_nodes": top_indices.tolist(), - "explanation_method": "gradient_norm" + "explanation_method": "gradient_norm", } diff --git a/drug_discovery/formulation/custom_excipient_matcher.py b/drug_discovery/formulation/custom_excipient_matcher.py index 788a965b4e..0fd67079dd 100644 --- a/drug_discovery/formulation/custom_excipient_matcher.py +++ b/drug_discovery/formulation/custom_excipient_matcher.py @@ -1,94 +1,107 @@ -import pandas as pd -import networkx as nx -from typing import List, Dict, Any, Optional import logging +from typing import Any + +import networkx as nx +import pandas as pd logger = logging.getLogger(__name__) + class PersonalizedFormulationEngine: """ - Matches drug candidates with biologically safe excipients tailored to + Matches drug candidates with biologically safe excipients tailored to specific patient health profiles and conditions. """ + def __init__(self): # Database of excipients and their contraindications # In a production system, this would be loaded from a validated medical database - self.excipient_knowledge_base = pd.DataFrame([ - {"name": "Sodium Chloride", "category": "Salt", "contraindications": ["hypernatremia", "hypertension"]}, - {"name": "Sucrose", "category": "Sweetener", "contraindications": ["diabetes", "fructose_intolerance"]}, - {"name": "Lactose", "category": "Filler", "contraindications": ["lactose_intolerance"]}, - {"name": "Ethanol", "category": "Solvent", "contraindications": ["liver_failure", "alcoholism", "pregnancy"]}, - {"name": "Propylene Glycol", "category": "Solvent", "contraindications": ["renal_failure"]}, - {"name": "Mannitol", "category": "Diuretic/Sweetener", "contraindications": ["anuria", "severe_dehydration"]}, - {"name": "Aspartame", "category": "Sweetener", "contraindications": ["phenylketonuria"]} - ]) - + self.excipient_knowledge_base = pd.DataFrame( + [ + {"name": "Sodium Chloride", "category": "Salt", "contraindications": ["hypernatremia", "hypertension"]}, + {"name": "Sucrose", "category": "Sweetener", "contraindications": ["diabetes", "fructose_intolerance"]}, + {"name": "Lactose", "category": "Filler", "contraindications": ["lactose_intolerance"]}, + { + "name": "Ethanol", + "category": "Solvent", + "contraindications": ["liver_failure", "alcoholism", "pregnancy"], + }, + {"name": "Propylene Glycol", "category": "Solvent", "contraindications": ["renal_failure"]}, + { + "name": "Mannitol", + "category": "Diuretic/Sweetener", + "contraindications": ["anuria", "severe_dehydration"], + }, + {"name": "Aspartame", "category": "Sweetener", "contraindications": ["phenylketonuria"]}, + ] + ) + # Build a relationship graph for cross-reactivity (optional complexity) self.reactivity_graph = nx.Graph() - def veto_contraindicated_vehicles(self, - patient_state: Any, - proposed_formulations: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + def veto_contraindicated_vehicles( + self, patient_state: Any, proposed_formulations: list[dict[str, Any]] + ) -> list[dict[str, Any]]: """ Filters out formulations that contain excipients incompatible with the patient's current health state. """ - patient_conditions = set(getattr(patient_state, 'conditions', [])) - + patient_conditions = set(getattr(patient_state, "conditions", [])) + # Derived conditions from biomarkers - sodium_level = getattr(patient_state, 'sodium_level', 140.0) - egfr = getattr(patient_state, 'egfr', 90.0) - ast = getattr(patient_state, 'ast', 25.0) - + sodium_level = getattr(patient_state, "sodium_level", 140.0) + egfr = getattr(patient_state, "egfr", 90.0) + ast = getattr(patient_state, "ast", 25.0) + if sodium_level > 145: patient_conditions.add("hypernatremia") if egfr < 30: patient_conditions.add("renal_failure") if ast > 120: patient_conditions.add("liver_failure") - + safe_formulations = [] - + for formulation in proposed_formulations: excipients = formulation.get("excipients", []) is_vetoed = False veto_reason = "" - + for excipient_name in excipients: # Lookup excipient in knowledge base match = self.excipient_knowledge_base[ - self.excipient_knowledge_base['name'].str.lower() == excipient_name.lower() + self.excipient_knowledge_base["name"].str.lower() == excipient_name.lower() ] - + if not match.empty: - contraindications = match.iloc[0]['contraindications'] + contraindications = match.iloc[0]["contraindications"] # Check for intersection between patient conditions and contraindications overlap = patient_conditions.intersection(set(contraindications)) if overlap: is_vetoed = True veto_reason = f"Excipient {excipient_name} is contraindicated for: {list(overlap)}" break - + if not is_vetoed: safe_formulations.append(formulation) else: logger.info(f"Vetoing formulation {formulation.get('id', 'unknown')}: {veto_reason}") - + return safe_formulations - def suggest_safe_alternatives(self, category: str, patient_state: Any) -> List[str]: + def suggest_safe_alternatives(self, category: str, patient_state: Any) -> list[str]: """ Suggests safe excipients within a specific category for a given patient. """ - patient_conditions = set(getattr(patient_state, 'conditions', [])) + patient_conditions = set(getattr(patient_state, "conditions", [])) # (Biomarker derived conditions logic here...) - + category_matches = self.excipient_knowledge_base[ - self.excipient_knowledge_base['category'].str.lower() == category.lower() + self.excipient_knowledge_base["category"].str.lower() == category.lower() ] - + alternatives = [] for _, row in category_matches.iterrows(): - if not set(row['contraindications']).intersection(patient_conditions): - alternatives.append(row['name']) - + if not set(row["contraindications"]).intersection(patient_conditions): + alternatives.append(row["name"]) + return alternatives diff --git a/drug_discovery/generation/analog_generator.py b/drug_discovery/generation/analog_generator.py index 603ac003fa..07d95d8fb4 100644 --- a/drug_discovery/generation/analog_generator.py +++ b/drug_discovery/generation/analog_generator.py @@ -1,9 +1,10 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any, Sequence + from rdkit import Chem + @dataclass class AnalogResult: parent_smiles: str @@ -11,13 +12,14 @@ class AnalogResult: tox_reduction: float = 0.0 potency_gain: float = 0.0 + class DeepGraphAnalogGenerator: """Deep graph networks for low-tox analog generation (2025). - + Generates analogs optimizing tox/potency. """ def generate_low_tox(self, parent: str, num_analogs: int = 10) -> AnalogResult: mol = Chem.MolFromSmiles(parent) analogs = [Chem.MolToSmiles(mol) for _ in range(num_analogs)] # mock mutate - return AnalogResult(parent, analogs, tox_reduction=0.3, potency_gain=1.2) \ No newline at end of file + return AnalogResult(parent, analogs, tox_reduction=0.3, potency_gain=1.2) diff --git a/drug_discovery/generation/backends.py b/drug_discovery/generation/backends.py index c951930724..0d40068232 100644 --- a/drug_discovery/generation/backends.py +++ b/drug_discovery/generation/backends.py @@ -261,7 +261,7 @@ def generate(self, prompt: str | None, num: int = 10, **kwargs) -> GenerationRes ) except Exception as e: - return GenerationResult.failure(self.name, f"NVIDIA LLM generation failed: {str(e)}") + return GenerationResult.failure(self.name, f"NVIDIA LLM generation failed: {e!s}") class GenerationManager: diff --git a/drug_discovery/generation/enhanced_retrosynth.py b/drug_discovery/generation/enhanced_retrosynth.py index cbc51ab17f..3eaa14a168 100644 --- a/drug_discovery/generation/enhanced_retrosynth.py +++ b/drug_discovery/generation/enhanced_retrosynth.py @@ -1,10 +1,10 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any, Sequence # Insilico/Exscientia style forward synthesis + retrosynth + @dataclass class SynthResult: target_smiles: str @@ -13,9 +13,10 @@ class SynthResult: accessibility_score: float = 0.0 success: bool = False + class EnhancedRetrosynth: """Enhanced retrosynthesis with forward synthesis planning (2024 generative breakthrough). - + Integrates rxnmapper + yield prediction. """ @@ -23,4 +24,4 @@ def plan_synthesis(self, target: str, max_paths: int = 5) -> SynthResult: paths = [["reactant1.reactant2", target] for _ in range(max_paths)] # mock yields = [0.85, 0.92, 0.78, 0.88, 0.91] score = sum(yields) / max_paths - return SynthResult(target, paths, yields, score, success=True) \ No newline at end of file + return SynthResult(target, paths, yields, score, success=True) diff --git a/drug_discovery/generation/physics_aware.py b/drug_discovery/generation/physics_aware.py index c3cf647fa7..8ab528b3e0 100644 --- a/drug_discovery/generation/physics_aware.py +++ b/drug_discovery/generation/physics_aware.py @@ -8,6 +8,7 @@ from __future__ import annotations +import contextlib import math import random from collections.abc import Sequence @@ -132,10 +133,8 @@ def _fragment_library(self, seeds: Sequence[str] | None) -> list[str]: mol = Chem.MolFromSmiles(smi) if mol is None: continue - try: + with contextlib.suppress(Exception): fragments.update(BRICS.BRICSDecompose(mol, silent=True)) - except Exception: - pass try: recap = rdMolDescriptors.RecapDecompose(mol) if recap: @@ -236,7 +235,7 @@ def _conformer_ensemble(self, smiles: str) -> ConformerEnsemble: t = 300.0 weights = [math.exp(-(e - min_e) / (k_b * t)) for e in energies] weight_sum = sum(weights) or 1.0 - boltz_avg = sum(e * w for e, w in zip(energies, weights)) / weight_sum + boltz_avg = sum(e * w for e, w in zip(energies, weights, strict=False)) / weight_sum ref_conf = int(energies.index(min_e)) ref_coords = mol.GetConformer(ref_conf).GetPositions() for idx, energy in enumerate(energies): diff --git a/drug_discovery/generation/torchdrug_generator.py b/drug_discovery/generation/torchdrug_generator.py index 9850c920fd..b82527242b 100644 --- a/drug_discovery/generation/torchdrug_generator.py +++ b/drug_discovery/generation/torchdrug_generator.py @@ -1,18 +1,22 @@ -import torch import logging + +import torch + try: import torchdrug as td from torchdrug import data, models, tasks, utils + TORCHDRUG_AVAILABLE = True except ImportError: TORCHDRUG_AVAILABLE = False td = None -from typing import List, Optional + from ..data.rdkit_utils import smiles_to_sdf logger = logging.getLogger(__name__) + class TorchDrugGenerator: def __init__(self, model_name: str = "VAE", num_layers: int = 3): self._fallback_mode = not TORCHDRUG_AVAILABLE @@ -31,14 +35,14 @@ def _build_model(self, model_name: str, num_layers: int): encoder_layers=num_layers, decoder_layers=num_layers, latent_size=128, - use_layer_norm=True + use_layer_norm=True, ) task = tasks.MoleculeGeneration(model, self.dataset, td.metrics.MoleculeMetrics("valid")) # Add GNN etc. task = task.to(self.device) return task - def generate(self, num: int = 1000, scaffold: Optional[str] = None) -> List[str]: + def generate(self, num: int = 1000, scaffold: str | None = None) -> list[str]: if self._fallback_mode or self.task is None: logger.warning("TorchDrug unavailable; returning placeholder molecules in fallback mode.") if scaffold and scaffold.strip(): @@ -51,5 +55,5 @@ def generate(self, num: int = 1000, scaffold: Optional[str] = None) -> List[str] smiles_list = [mol.smiles() for mol in generated] return smiles_list - def save_to_sdf(self, smiles_list: List[str], path: str): + def save_to_sdf(self, smiles_list: list[str], path: str): smiles_to_sdf(smiles_list, path) diff --git a/drug_discovery/geometric_dl/__init__.py b/drug_discovery/geometric_dl/__init__.py index bdbef73925..428934f969 100644 --- a/drug_discovery/geometric_dl/__init__.py +++ b/drug_discovery/geometric_dl/__init__.py @@ -31,14 +31,14 @@ ) __all__ = [ - "SE3Transformer", - "SE3Config", - "SE3EquivariantBlock", - "EquivariantAttention", "BindingFreeEnergyCalculator", + "EquivariantAttention", "FEPConfig", "FEPResult", "OpenMMDriver", - "TransientPocketPredictor", "PocketPrediction", + "SE3Config", + "SE3EquivariantBlock", + "SE3Transformer", + "TransientPocketPredictor", ] diff --git a/drug_discovery/geometric_dl/fep_engine.py b/drug_discovery/geometric_dl/fep_engine.py index 70c26935d7..db86ea4c57 100644 --- a/drug_discovery/geometric_dl/fep_engine.py +++ b/drug_discovery/geometric_dl/fep_engine.py @@ -155,10 +155,7 @@ def _setup_platform(self) -> None: try: self.platform = Platform.getPlatformByName(self.platform_name) - if self.platform_name == "CUDA": - props = {"DeviceIndex": str(self.device_index), "Precision": self.precision} - self.platform.setPropertyDefaultValues(props) - elif self.platform_name == "OpenCL": + if self.platform_name == "CUDA" or self.platform_name == "OpenCL": props = {"DeviceIndex": str(self.device_index), "Precision": self.precision} self.platform.setPropertyDefaultValues(props) diff --git a/drug_discovery/geometric_dl/pocket_predictor.py b/drug_discovery/geometric_dl/pocket_predictor.py index 7299c1fef5..b6ab049676 100644 --- a/drug_discovery/geometric_dl/pocket_predictor.py +++ b/drug_discovery/geometric_dl/pocket_predictor.py @@ -173,7 +173,7 @@ def _compute_surface( surface_points = [] # Simple approximation: points at van der Waals radius from atoms - for i, (pos, atype) in enumerate(zip(coords, atom_types)): + for _i, (pos, atype) in enumerate(zip(coords, atom_types, strict=False)): # Approximate radius based on atom type radius = self._vdw_radius(atype) @@ -182,7 +182,7 @@ def _compute_surface( theta = np.random.uniform(0, 2 * np.pi, n_samples) phi = np.random.uniform(0, np.pi, n_samples) - for th, ph in zip(theta, phi): + for th, ph in zip(theta, phi, strict=False): point = pos + radius * np.array( [ np.sin(ph) * np.cos(th), @@ -219,7 +219,7 @@ def _is_surface_point( atom_types: np.ndarray, ) -> bool: """Check if point is on surface (not buried).""" - for i, (pos, atype) in enumerate(zip(coords, atom_types)): + for _i, (pos, atype) in enumerate(zip(coords, atom_types, strict=False)): dist = np.linalg.norm(point - pos) radius = self._vdw_radius(atype) @@ -308,10 +308,9 @@ def _find_nearby_residues( ) -> list[int]: """Find residues within cutoff of pocket center.""" nearby = [] - for i, (pos, res_idx) in enumerate(zip(protein_coords, residue_indices)): - if np.linalg.norm(pos - center) < cutoff: - if res_idx not in nearby: - nearby.append(int(res_idx)) + for _i, (pos, res_idx) in enumerate(zip(protein_coords, residue_indices, strict=False)): + if np.linalg.norm(pos - center) < cutoff and res_idx not in nearby: + nearby.append(int(res_idx)) return nearby def _score_pockets( diff --git a/drug_discovery/geometric_dl/se3_transformer.py b/drug_discovery/geometric_dl/se3_transformer.py index de8d8de9d2..3b5175a716 100644 --- a/drug_discovery/geometric_dl/se3_transformer.py +++ b/drug_discovery/geometric_dl/se3_transformer.py @@ -144,7 +144,7 @@ def forward(self, x, edge_index, edge_attr, edge_sh): if not E3NN_AVAILABLE: return x - src, dst = edge_index + src, _dst = edge_index # Message: tensor product of input features with spherical harmonics messages = self.tp(x[src], edge_sh) @@ -309,7 +309,7 @@ def __init__(self, config: SE3Config | None = None): # SE(3)-Equivariant layers self.layers = nn.ModuleList() - for i in range(self.config.num_layers): + for _i in range(self.config.num_layers): self.layers.append( SE3EquivariantBlock( irreps_in=f"{hidden_dim}x0e", diff --git a/drug_discovery/glp_tox_panel.py b/drug_discovery/glp_tox_panel.py index dda8e2341f..5d293bb077 100644 --- a/drug_discovery/glp_tox_panel.py +++ b/drug_discovery/glp_tox_panel.py @@ -26,6 +26,7 @@ try: from drug_discovery.evaluation.herg_predictor import HERGPredictor + _HERG_PREDICTOR = True except ImportError: _HERG_PREDICTOR = False @@ -164,7 +165,7 @@ def __init__( self.herg_threshold = herg_threshold self.cyp_threshold = cyp_threshold self.ames_threshold = ames_threshold - + # Use provided predictor or create default if herg_predictor is not None: self.herg_predictor = herg_predictor @@ -189,9 +190,11 @@ def evaluate(self, smiles: str) -> GLPToxPanel: if not ames.passed: rejection_reasons.append(f"Ames: {ames.risk_class}") - overall = (herg.inhibition_probability * 0.4 - + max(cyp.enzyme_inhibitions.values(), default=0) * 0.3 - + ames.mutagenicity_probability * 0.3) + overall = ( + herg.inhibition_probability * 0.4 + + max(cyp.enzyme_inhibitions.values(), default=0) * 0.3 + + ames.mutagenicity_probability * 0.3 + ) return GLPToxPanel( smiles=smiles, @@ -218,14 +221,14 @@ def _virtual_herg(self, smiles: str, props: dict[str, float]) -> HERGResult: # Use new HERG predictor if available if self.herg_predictor is not None and _HERG_PREDICTOR: prediction = self.herg_predictor.predict(smiles, calibrate=True) - + # Map to CiPA risk classification risk_class = "low" if prediction.cipa_risk_category == "category_2": risk_class = "moderate" elif prediction.cipa_risk_category == "category_3": risk_class = "high" - + return HERGResult( smiles=smiles, inhibition_probability=prediction.inhibition_probability, @@ -234,7 +237,7 @@ def _virtual_herg(self, smiles: str, props: dict[str, float]) -> HERGResult: passed=prediction.inhibition_probability <= self.herg_threshold, key_features=prediction.key_concerns, ) - + # Fallback: original heuristic model when new predictor unavailable logp = props.get("logp", 2.0) tpsa = props.get("tpsa", 60.0) @@ -248,8 +251,7 @@ def _virtual_herg(self, smiles: str, props: dict[str, float]) -> HERGResult: basicity_factor = _sigmoid(hbd - 2) # Weighted combination (normalized) - prob = (logp_factor * 0.40 + tpsa_factor * 0.25 + - mw_factor * 0.20 + basicity_factor * 0.15) + prob = logp_factor * 0.40 + tpsa_factor * 0.25 + mw_factor * 0.20 + basicity_factor * 0.15 prob = min(max(prob, 0.0), 1.0) if prob > 0.7: diff --git a/drug_discovery/intelligence/__init__.py b/drug_discovery/intelligence/__init__.py index 2dd71efd0a..f0eb6c6b27 100644 --- a/drug_discovery/intelligence/__init__.py +++ b/drug_discovery/intelligence/__init__.py @@ -28,13 +28,13 @@ from .rag_engine import RAGEngine as RAGEngine __all__ = [ + "ArXivIngester", "BiomedicalIntelligence", "BiomedicalNER", - "RelationshipExtractor", - "PubMedIngester", - "ArXivIngester", - "LiteratureDocument", "ExtractedEntity", "ExtractedRelationship", + "LiteratureDocument", + "PubMedIngester", "RAGEngine", + "RelationshipExtractor", ] diff --git a/drug_discovery/knowledge_graph/__init__.py b/drug_discovery/knowledge_graph/__init__.py index 1e25d4ea60..42fab67fdc 100644 --- a/drug_discovery/knowledge_graph/__init__.py +++ b/drug_discovery/knowledge_graph/__init__.py @@ -23,21 +23,21 @@ ) from .link_prediction import LinkPredictionService as LinkPredictionService from .link_prediction import LinkPredictorGNN as LinkPredictorGNN -from .reasoning import KnowledgeGraphReasoner as KnowledgeGraphReasoner from .neo4j_adapter import Neo4jAdapter as Neo4jAdapter +from .reasoning import KnowledgeGraphReasoner as KnowledgeGraphReasoner __all__ = [ "DrugKnowledgeGraph", - "KnowledgeGraphBuilder", - "KnowledgeGraph", - "VectorDatabase", - "KGNode", - "KGEdge", - "NodeType", "EdgeType", - "Neo4jAdapter", - "LinkPredictorGNN", - "LinkPredictionService", - "KnowledgeGraphReasoner", + "KGEdge", "KGIngestor", + "KGNode", + "KnowledgeGraph", + "KnowledgeGraphBuilder", + "KnowledgeGraphReasoner", + "LinkPredictionService", + "LinkPredictorGNN", + "Neo4jAdapter", + "NodeType", + "VectorDatabase", ] diff --git a/drug_discovery/knowledge_graph/graph.py b/drug_discovery/knowledge_graph/graph.py index 06c8b3806e..bffabf008a 100644 --- a/drug_discovery/knowledge_graph/graph.py +++ b/drug_discovery/knowledge_graph/graph.py @@ -79,14 +79,14 @@ def query_relations( if direction in ["outgoing", "both"]: for target in self.graph.successors(entity_id): edges = self.graph[entity_id][target] - for key, data in edges.items(): + for _key, data in edges.items(): if relation_type is None or data.get("relation") == relation_type: results.append((entity_id, target, data)) if direction in ["incoming", "both"]: for source in self.graph.predecessors(entity_id): edges = self.graph[source][entity_id] - for key, data in edges.items(): + for _key, data in edges.items(): if relation_type is None or data.get("relation") == relation_type: results.append((source, entity_id, data)) diff --git a/drug_discovery/knowledge_graph/ingestion.py b/drug_discovery/knowledge_graph/ingestion.py index b2fea7fc4b..d7e57aa900 100644 --- a/drug_discovery/knowledge_graph/ingestion.py +++ b/drug_discovery/knowledge_graph/ingestion.py @@ -30,7 +30,7 @@ def ingest_batch(batch_data: list[dict[str, any]], node_type: str): class KGIngestor: """Orchestrates large-scale ingestion of data into the knowledge graph.""" - def __init__(self, batch_size: int = 1000, num_workers: int = None): + def __init__(self, batch_size: int = 1000, num_workers: int | None = None): self.batch_size = batch_size self.num_workers = num_workers or multiprocessing.cpu_count() diff --git a/drug_discovery/knowledge_graph/knowledge_graph.py b/drug_discovery/knowledge_graph/knowledge_graph.py index f0b323893b..e1b09d89b4 100644 --- a/drug_discovery/knowledge_graph/knowledge_graph.py +++ b/drug_discovery/knowledge_graph/knowledge_graph.py @@ -270,10 +270,7 @@ def get_neighbors( continue # Get neighbor node - if edge.source_id == node_id: - neighbor_id = edge.target_id - else: - neighbor_id = edge.source_id + neighbor_id = edge.target_id if edge.source_id == node_id else edge.source_id if neighbor_id in self.nodes: neighbors.append(self.nodes[neighbor_id]) @@ -320,7 +317,7 @@ def find_path( if neighbor_id not in visited: visited.add(neighbor_id) - new_path = path + [(neighbor_id, edge_id)] + new_path = [*path, (neighbor_id, edge_id)] queue.append((neighbor_id, new_path)) return None @@ -469,7 +466,7 @@ def get_subgraph( edges = [] if include_edges: node_id_set = set(node_ids) - for edge_id, edge in self.edges.items(): + for _edge_id, edge in self.edges.items(): if edge.source_id in node_id_set and edge.target_id in node_id_set: edges.append(edge) diff --git a/drug_discovery/knowledge_graph/link_prediction.py b/drug_discovery/knowledge_graph/link_prediction.py index e9d9a6498a..fc71b99a55 100644 --- a/drug_discovery/knowledge_graph/link_prediction.py +++ b/drug_discovery/knowledge_graph/link_prediction.py @@ -9,7 +9,6 @@ import torch import torch.nn as nn import torch.nn.functional as F - from torch_geometric.data import Data from torch_geometric.nn import GATConv diff --git a/drug_discovery/knowledge_graph/reasoning.py b/drug_discovery/knowledge_graph/reasoning.py index 672c039a8a..ef7797562e 100644 --- a/drug_discovery/knowledge_graph/reasoning.py +++ b/drug_discovery/knowledge_graph/reasoning.py @@ -1,16 +1,18 @@ -from typing import List, Dict, Tuple, Set, Any import networkx as nx + class KnowledgeGraphReasoner: """ Symbolic and path-based reasoning for Drug Discovery Knowledge Graphs. Can answer queries like "Find all proteins regulated by genes associated with disease X". """ - + def __init__(self, nx_graph: nx.MultiDiGraph): self.graph = nx_graph - - def find_causal_paths(self, start_node: str, end_node: str, max_length: int = 3) -> List[List[Tuple[str, str, str]]]: + + def find_causal_paths( + self, start_node: str, end_node: str, max_length: int = 3 + ) -> list[list[tuple[str, str, str]]]: """ Find all directed paths between two nodes up to a certain length. Returns paths as list of (u, v, edge_type) tuples. @@ -18,25 +20,25 @@ def find_causal_paths(self, start_node: str, end_node: str, max_length: int = 3) paths = [] if start_node not in self.graph or end_node not in self.graph: return [] - + for path in nx.all_simple_paths(self.graph, start_node, end_node, cutoff=max_length): path_with_types = [] for i in range(len(path) - 1): - u, v = path[i], path[i+1] + u, v = path[i], path[i + 1] # Get the edge type (assuming it's stored in 'edge_type' attribute) edge_data = self.graph.get_edge_data(u, v) # MultiDiGraph might have multiple edges between same nodes etype = "unknown" if edge_data: # Just take the first one for simplicity - key = list(edge_data.keys())[0] - etype = edge_data[key].get('edge_type', 'unknown') + key = next(iter(edge_data.keys())) + etype = edge_data[key].get("edge_type", "unknown") path_with_types.append((u, v, etype)) paths.append(path_with_types) - + return paths - def query_by_relation_chain(self, start_node: str, relations: List[str]) -> Set[str]: + def query_by_relation_chain(self, start_node: str, relations: list[str]) -> set[str]: """ Multi-hop query: Start at start_node, follow the sequence of relations. Example: relations=['ASSOCIATED_WITH_GENE', 'CODES_FOR_PROTEIN'] @@ -48,26 +50,26 @@ def query_by_relation_chain(self, start_node: str, relations: List[str]) -> Set[ if node not in self.graph: continue for _, neighbor, data in self.graph.out_edges(node, data=True): - if data.get('edge_type') == rel: + if data.get("edge_type") == rel: next_nodes.add(neighbor) current_nodes = next_nodes if not current_nodes: break return current_nodes - def check_consistency(self) -> List[str]: + def check_consistency(self) -> list[str]: """ Check for logical inconsistencies in the KG. Example: A drug cannot both inhibit and activate the same protein simultaneously. """ issues = [] for u, v, data in self.graph.edges(data=True): - etype = data.get('edge_type') - if etype == 'INHIBITS': + etype = data.get("edge_type") + if etype == "INHIBITS": # Check if there's also an 'ACTIVATES' edge if self.graph.has_edge(u, v): edge_data = self.graph.get_edge_data(u, v) for key in edge_data: - if edge_data[key].get('edge_type') == 'ACTIVATES': + if edge_data[key].get("edge_type") == "ACTIVATES": issues.append(f"Contradictory relationship between {u} and {v}: INHIBITS and ACTIVATES") return issues diff --git a/drug_discovery/lead_optimization/__init__.py b/drug_discovery/lead_optimization/__init__.py index 8b38b05312..7f7578868b 100644 --- a/drug_discovery/lead_optimization/__init__.py +++ b/drug_discovery/lead_optimization/__init__.py @@ -3,8 +3,8 @@ Methods for optimizing drug candidates using MCTS and RL. """ +from .ensemble_refiner import EnsembleRefiner from .mcts import LeadMCTSOptimizer from .rl_optimizer import LeadRLOptimizer -from .ensemble_refiner import EnsembleRefiner -__all__ = ["LeadMCTSOptimizer", "LeadRLOptimizer", "EnsembleRefiner"] +__all__ = ["EnsembleRefiner", "LeadMCTSOptimizer", "LeadRLOptimizer"] diff --git a/drug_discovery/lead_optimization/ensemble_refiner.py b/drug_discovery/lead_optimization/ensemble_refiner.py index 3d1098e45e..c54b4a96d3 100644 --- a/drug_discovery/lead_optimization/ensemble_refiner.py +++ b/drug_discovery/lead_optimization/ensemble_refiner.py @@ -1,7 +1,8 @@ -from typing import List, Dict, Any, Optional -import torch +from typing import Any + import numpy as np + class EnsembleRefiner: """ Refines drug leads by combining multiple scoring functions: @@ -10,116 +11,105 @@ class EnsembleRefiner: - Physics-based binding affinity (optional) - Synthetic Accessibility (SA) """ - + def __init__( self, property_predictor: Any, admet_predictor: Any, - physics_simulator: Optional[Any] = None, - weights: Optional[Dict[str, float]] = None + physics_simulator: Any | None = None, + weights: dict[str, float] | None = None, ): self.property_predictor = property_predictor self.admet_predictor = admet_predictor self.physics_simulator = physics_simulator - self.weights = weights or { - 'property': 0.4, - 'admet': 0.3, - 'physics': 0.2, - 'sa': 0.1 - } - - def calculate_ensemble_score(self, smiles: str, target_protein_pdb: Optional[str] = None) -> float: + self.weights = weights or {"property": 0.4, "admet": 0.3, "physics": 0.2, "sa": 0.1} + + def calculate_ensemble_score(self, smiles: str, target_protein_pdb: str | None = None) -> float: """ Calculate a weighted aggregate score for a molecule with multi-objective optimization. """ scores = {} - + # 1. Property Score (Activity/Potency) try: prop_res = self.property_predictor.predict(smiles) # If it's a dict, take the value, else it's the score itself - scores['property'] = prop_res if isinstance(prop_res, (int, float)) else list(prop_res.values())[0] + scores["property"] = prop_res if isinstance(prop_res, (int, float)) else next(iter(prop_res.values())) except Exception: - scores['property'] = 0.0 - + scores["property"] = 0.0 + # 2. ADMET Score (Safety and Pharmokinetics) try: # We want high QED and high safety qed = self.admet_predictor.calculate_qed(smiles) - tox_verdict = self.admet_predictor.evaluate(smiles) if hasattr(self.admet_predictor, 'evaluate') else None + tox_verdict = self.admet_predictor.evaluate(smiles) if hasattr(self.admet_predictor, "evaluate") else None safety = tox_verdict.safety_score if tox_verdict else 1.0 - scores['admet'] = 0.5 * qed + 0.5 * safety + scores["admet"] = 0.5 * qed + 0.5 * safety except Exception: - scores['admet'] = 0.0 - + scores["admet"] = 0.0 + # 3. SA Score (Synthetic Accessibility) try: sa = self.admet_predictor.calculate_synthetic_accessibility(smiles) # 1.0 is very easy, 0.0 is impossible - scores['sa'] = 1.0 - (sa / 10.0) + scores["sa"] = 1.0 - (sa / 10.0) except Exception: - scores['sa'] = 0.0 - + scores["sa"] = 0.0 + # 4. Physics Score (Binding Affinity) if self.physics_simulator and target_protein_pdb: try: res = self.physics_simulator.simulate_md(smiles, target_protein_pdb) - binding_energy = res.get('binding_energy', 0.0) + binding_energy = res.get("binding_energy", 0.0) # Use a sigmoid to normalize energy: -10 kcal/mol should be a very high score - scores['physics'] = 1.0 / (1.0 + np.exp(0.5 * (binding_energy + 8.0))) + scores["physics"] = 1.0 / (1.0 + np.exp(0.5 * (binding_energy + 8.0))) except Exception: - scores['physics'] = 0.0 + scores["physics"] = 0.0 else: - scores['physics'] = 0.0 + scores["physics"] = 0.0 # Multi-objective Pareto-like aggregate (Geometric Mean for balance) # We add a small epsilon to avoid zeroing out the entire score epsilon = 0.01 weighted_scores = [ - (scores.get('property', 0) + epsilon) ** self.weights.get('property', 0.4), - (scores.get('admet', 0) + epsilon) ** self.weights.get('admet', 0.3), - (scores.get('physics', 0) + epsilon) ** self.weights.get('physics', 0.2), - (scores.get('sa', 0) + epsilon) ** self.weights.get('sa', 0.1) + (scores.get("property", 0) + epsilon) ** self.weights.get("property", 0.4), + (scores.get("admet", 0) + epsilon) ** self.weights.get("admet", 0.3), + (scores.get("physics", 0) + epsilon) ** self.weights.get("physics", 0.2), + (scores.get("sa", 0) + epsilon) ** self.weights.get("sa", 0.1), ] - + total_score = np.prod(weighted_scores) return float(total_score) def select_elite_candidates( - self, - smiles_list: List[str], - target_protein_pdb: Optional[str] = None, - top_k: int = 10 - ) -> List[Dict[str, Any]]: + self, smiles_list: list[str], target_protein_pdb: str | None = None, top_k: int = 10 + ) -> list[dict[str, Any]]: """ Select 'Elite' candidates that excel across all objectives. """ results = self.rank_candidates(smiles_list, target_protein_pdb) - + # Additional filter for elite candidates: must pass safety gate strictly elite = [] for res in results: if len(elite) >= top_k: break - + try: # Assuming admet_predictor has a validate method that checks elite standards - if hasattr(self.admet_predictor, 'is_elite_smiles'): - if not self.admet_predictor.is_elite_smiles(res['smiles']): + if hasattr(self.admet_predictor, "is_elite_smiles"): + if not self.admet_predictor.is_elite_smiles(res["smiles"]): continue except Exception: pass - + elite.append(res) - + return elite - def rank_candidates(self, smiles_list: List[str], target_protein_pdb: Optional[str] = None) -> List[Dict[str, Any]]: + def rank_candidates(self, smiles_list: list[str], target_protein_pdb: str | None = None) -> list[dict[str, Any]]: results = [] for smiles in smiles_list: score = self.calculate_ensemble_score(smiles, target_protein_pdb) - results.append({ - 'smiles': smiles, - 'ensemble_score': score - }) - return sorted(results, key=lambda x: x['ensemble_score'], reverse=True) + results.append({"smiles": smiles, "ensemble_score": score}) + return sorted(results, key=lambda x: x["ensemble_score"], reverse=True) diff --git a/drug_discovery/lead_optimization/mcts.py b/drug_discovery/lead_optimization/mcts.py index b7fde562d7..d438e4f944 100644 --- a/drug_discovery/lead_optimization/mcts.py +++ b/drug_discovery/lead_optimization/mcts.py @@ -1,16 +1,18 @@ +from typing import Any + import numpy as np -from typing import List, Dict, Any, Optional + class LeadMCTSOptimizer: """ Lead Optimization using Monte Carlo Tree Search. Explores the chemical space by adding/removing/replacing fragments. """ - - def __init__(self, property_predictor: Any, fragments: List[str]): + + def __init__(self, property_predictor: Any, fragments: list[str]): self.predictor = property_predictor self.fragments = fragments - + def optimize(self, seed_smiles: str, iterations: int = 100) -> str: """ Search for an optimized version of the seed molecule. @@ -18,20 +20,20 @@ def optimize(self, seed_smiles: str, iterations: int = 100) -> str: current_smiles = seed_smiles best_smiles = seed_smiles best_score = self.predictor.predict(seed_smiles) - + for _ in range(iterations): # 1. Selection & Expansion (Simplified) new_smiles = self._mutate(current_smiles) - + # 2. Simulation (Evaluation) score = self.predictor.predict(new_smiles) - + # 3. Backpropagation (Update Best) if score > best_score: best_score = score best_smiles = new_smiles - current_smiles = new_smiles # Greedy exploration - + current_smiles = new_smiles # Greedy exploration + return best_smiles def _mutate(self, smiles: str) -> str: diff --git a/drug_discovery/lead_optimization/rl_optimizer.py b/drug_discovery/lead_optimization/rl_optimizer.py index f775eea3dd..8966a5f432 100644 --- a/drug_discovery/lead_optimization/rl_optimizer.py +++ b/drug_discovery/lead_optimization/rl_optimizer.py @@ -1,39 +1,41 @@ +from typing import Any + import torch import torch.nn as nn import torch.optim as optim -from typing import List, Any + class LeadRLOptimizer: """ Lead Optimization using Reinforcement Learning (Policy Gradient). """ - + def __init__(self, policy_network: nn.Module, property_predictor: Any): self.policy = policy_network self.predictor = property_predictor self.optimizer = optim.Adam(self.policy.parameters(), lr=1e-3) - - def train_step(self, states: torch.Tensor, actions: torch.Tensor, smiles_list: List[str]): + + def train_step(self, states: torch.Tensor, actions: torch.Tensor, smiles_list: list[str]): """ Single RL training step. """ self.optimizer.zero_grad() - + # Calculate rewards based on predicted properties rewards = [] for smiles in smiles_list: reward = self.predictor.predict(smiles) rewards.append(reward) - + rewards = torch.tensor(rewards, dtype=torch.float32) # Normalize rewards rewards = (rewards - rewards.mean()) / (rewards.std() + 1e-9) - + # Policy gradient update (Simplified) log_probs = self.policy(states).gather(1, actions.unsqueeze(1)) loss = -(log_probs * rewards.unsqueeze(1)).mean() - + loss.backward() self.optimizer.step() - + return loss.item() diff --git a/drug_discovery/meta_learning/__init__.py b/drug_discovery/meta_learning/__init__.py index fd4f3b6fa6..8ceda41749 100644 --- a/drug_discovery/meta_learning/__init__.py +++ b/drug_discovery/meta_learning/__init__.py @@ -6,4 +6,4 @@ SelfImprovementOrchestrator, ) -__all__ = ["HypothesisGenerator", "CodeMutator", "SelfImprovementOrchestrator"] +__all__ = ["CodeMutator", "HypothesisGenerator", "SelfImprovementOrchestrator"] diff --git a/drug_discovery/models/__init__.py b/drug_discovery/models/__init__.py index 49cb0e419c..c27f812cc1 100644 --- a/drug_discovery/models/__init__.py +++ b/drug_discovery/models/__init__.py @@ -78,9 +78,9 @@ pass try: + from drug_discovery.models.gnn import MolecularGIN as MolecularGIN from drug_discovery.models.gnn import MolecularGNN as MolecularGNN from drug_discovery.models.gnn import MolecularMPNN as MolecularMPNN - from drug_discovery.models.gnn import MolecularGIN as MolecularGIN from drug_discovery.models.hetero_gnn import HeteroGNN as HeteroGNN MODEL_REGISTRY["gnn"] = {"class": MolecularGNN, "config": None, "variant": None} @@ -90,8 +90,8 @@ pass try: - from drug_discovery.models.transformer import MolecularTransformer as MolecularTransformer from drug_discovery.models.transformer import ModernMolecularTransformer as ModernMolecularTransformer + from drug_discovery.models.transformer import MolecularTransformer as MolecularTransformer MODEL_REGISTRY["transformer"] = { "class": MolecularTransformer, @@ -121,27 +121,27 @@ __all__ = [ "MODEL_REGISTRY", - "MolecularGNN", - "MolecularMPNN", - "MolecularGIN", - "HeteroGNN", - "MolecularTransformer", - "EnsembleModel", - "HybridModel", - "MultiTaskModel", - "EquivariantGNN", - "EquivariantGNNConfig", - "EGNNLayer", - "SchNetLayer", - "GaussianRBF", "CosineCutoff", - "build_radius_graph", - "MolecularDiffusionModel", "DiffusionConfig", "DiffusionMoleculeGenerator", - "GFlowNetPolicy", + "EGNNLayer", + "EnsembleModel", + "EquivariantGNN", + "EquivariantGNNConfig", "GFlowNetConfig", + "GFlowNetPolicy", "GFlowNetTrainer", - "PhysicsRewardFunction", + "GaussianRBF", + "HeteroGNN", + "HybridModel", + "MolecularDiffusionModel", + "MolecularGIN", + "MolecularGNN", + "MolecularMPNN", + "MolecularTransformer", + "MultiTaskModel", "PhysicsRewardConfig", + "PhysicsRewardFunction", + "SchNetLayer", + "build_radius_graph", ] diff --git a/drug_discovery/models/diffusion_generator.py b/drug_discovery/models/diffusion_generator.py index dd37240e24..e188897ca9 100644 --- a/drug_discovery/models/diffusion_generator.py +++ b/drug_discovery/models/diffusion_generator.py @@ -134,10 +134,7 @@ def sample(self, num_molecules, num_atoms, edge_index_fn=None): batch = torch.arange(num_molecules, device=self.device).repeat_interleave(num_atoms) for t_val in reversed(range(self.config.noise_steps)): t = torch.full((num_molecules,), t_val, device=self.device, dtype=torch.long) - if edge_index_fn: - edge_index = edge_index_fn(pos, batch) - else: - edge_index = self._fully_connected(num_molecules, num_atoms) + edge_index = edge_index_fn(pos, batch) if edge_index_fn else self._fully_connected(num_molecules, num_atoms) eps_pos, eps_atom = self.model(atom_types, pos, edge_index, t, batch) ab = self.alpha_bar[t_val] beta = 1 - ab / (self.alpha_bar[t_val - 1] if t_val > 0 else 1.0) diff --git a/drug_discovery/models/e3_equivariant.py b/drug_discovery/models/e3_equivariant.py index 515190994a..dc259d8d3d 100644 --- a/drug_discovery/models/e3_equivariant.py +++ b/drug_discovery/models/e3_equivariant.py @@ -8,7 +8,6 @@ import torch import torch.nn as nn import torch.nn.functional as F - from torch_geometric.nn import MessagePassing, global_mean_pool logger = logging.getLogger(__name__) @@ -98,10 +97,7 @@ def forward(self, data): x_new = F.relu(x_new) # Residual connection - if i > 0: - x = x + x_new - else: - x = x_new + x = x + x_new if i > 0 else x_new # Global pooling x = global_mean_pool(x, batch) diff --git a/drug_discovery/models/gnn.py b/drug_discovery/models/gnn.py index 2102dabc59..3ed273001e 100644 --- a/drug_discovery/models/gnn.py +++ b/drug_discovery/models/gnn.py @@ -5,7 +5,6 @@ import torch import torch.nn as nn import torch.nn.functional as F - from torch_geometric.nn import GATConv, GINConv, MessagePassing, global_max_pool, global_mean_pool @@ -100,10 +99,7 @@ def forward(self, data): x_new = self.dropout(x_new) # Residual connection - if i > 0: - x = x + x_new - else: - x = x_new + x = x + x_new if i > 0 else x_new # Graph pooling if self.pooling == "mean": diff --git a/drug_discovery/models/hetero_gnn.py b/drug_discovery/models/hetero_gnn.py index 1b8d1b26c3..8420f69111 100644 --- a/drug_discovery/models/hetero_gnn.py +++ b/drug_discovery/models/hetero_gnn.py @@ -1,34 +1,31 @@ import torch import torch.nn as nn -import torch.nn.functional as F -from torch_geometric.nn import HeteroConv, GATConv, Linear -from typing import Dict, List, Optional, Tuple, Any +from torch_geometric.nn import GATConv, HeteroConv, Linear + class HeteroGNN(nn.Module): """ Heterogeneous Graph Neural Network for multi-relational drug discovery. Supports DRUG, PROTEIN, GENE, DISEASE node types and their interactions. """ - + def __init__( self, - node_types: List[str], - edge_types: List[Tuple[str, str, str]], + node_types: list[str], + edge_types: list[tuple[str, str, str]], hidden_dim: int = 128, num_layers: int = 3, output_dim: int = 1, - dropout: float = 0.2 + dropout: float = 0.2, ): super().__init__() - + self.node_types = node_types self.edge_types = edge_types - + # Node encoders for each type - self.node_encoders = nn.ModuleDict({ - node_type: Linear(-1, hidden_dim) for node_type in node_types - }) - + self.node_encoders = nn.ModuleDict({node_type: Linear(-1, hidden_dim) for node_type in node_types}) + # Heterogeneous convolution layers self.convs = nn.ModuleList() for _ in range(num_layers): @@ -36,9 +33,9 @@ def __init__( for edge_type in edge_types: # edge_type is (src, rel, dst) conv_dict[edge_type] = GATConv((-1, -1), hidden_dim, add_self_loops=False) - - self.convs.append(HeteroConv(conv_dict, aggr='sum')) - + + self.convs.append(HeteroConv(conv_dict, aggr="sum")) + # Final prediction heads self.fc1 = nn.Linear(hidden_dim, hidden_dim // 2) self.fc2 = nn.Linear(hidden_dim // 2, output_dim) @@ -48,28 +45,28 @@ def forward(self, x_dict, edge_index_dict): # Encode nodes for node_type, x in x_dict.items(): x_dict[node_type] = self.node_encoders[node_type](x).relu() - + # Multi-layer heterogeneous convolutions for conv in self.convs: x_dict = conv(x_dict, edge_index_dict) - for node_type in x_dict.keys(): + for node_type in x_dict: x_dict[node_type] = x_dict[node_type].relu() x_dict[node_type] = self.dropout(x_dict[node_type]) - + # Return hidden representations (can be used for further tasks) return x_dict - def predict_link(self, h_dict: Dict[str, torch.Tensor], edge_type: Tuple[str, str, str], edge_index: torch.Tensor): + def predict_link(self, h_dict: dict[str, torch.Tensor], edge_type: tuple[str, str, str], edge_index: torch.Tensor): """ Predict link presence or property between two nodes. """ src_type, _, dst_type = edge_type src_h = h_dict[src_type][edge_index[0]] dst_h = h_dict[dst_type][edge_index[1]] - + # Combine representations h = src_h * dst_h - + h = self.fc1(h).relu() h = self.dropout(h) return self.fc2(h) diff --git a/drug_discovery/models/transformer.py b/drug_discovery/models/transformer.py index 56b5611e52..470753f59d 100644 --- a/drug_discovery/models/transformer.py +++ b/drug_discovery/models/transformer.py @@ -2,13 +2,32 @@ Transformer-based Models for Molecular Property Prediction """ -from typing import cast +import math import torch import torch.nn as nn import torch.nn.functional as F +class PositionalEncoding(nn.Module): + """Standard sinusoidal positional encoding for transformer models.""" + + def __init__(self, d_model: int, dropout: float = 0.1, max_len: int = 512): + super().__init__() + self.dropout = nn.Dropout(p=dropout) + pe = torch.zeros(max_len, d_model) + position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) + div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) + pe[:, 0::2] = torch.sin(position * div_term) + pe[:, 1::2] = torch.cos(position * div_term[: d_model // 2]) + pe = pe.unsqueeze(0) + self.register_buffer("pe", pe) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x + self.pe[:, : x.size(1)] + return self.dropout(x) + + class MolecularTransformer(nn.Module): """ Transformer model for molecular fingerprints or descriptors @@ -100,6 +119,7 @@ class SwiGLU(nn.Module): SwiGLU activation function (used in LLaMA and other modern transformers). Ref: Shazeer, "GLU Variants Improve Transformer" (2020). """ + def __init__(self, input_dim: int, output_dim: int): super().__init__() self.w1 = nn.Linear(input_dim, output_dim) @@ -113,6 +133,7 @@ class ModernMolecularTransformer(nn.Module): """ Improved Transformer model using SwiGLU activations and Pre-Norm. """ + def __init__( self, input_dim: int = 2048, @@ -124,15 +145,15 @@ def __init__( ): super().__init__() self.input_projection = nn.Linear(input_dim, hidden_dim) - + encoder_layer = nn.TransformerEncoderLayer( d_model=hidden_dim, nhead=num_heads, dim_feedforward=hidden_dim * 4, dropout=dropout, - activation=lambda x: F.silu(x), # SwiGLU is used in the FFN part, but PyTorch's layer is limited + activation=lambda x: F.silu(x), # SwiGLU is used in the FFN part, but PyTorch's layer is limited batch_first=True, - norm_first=True # Pre-Norm for better stability + norm_first=True, # Pre-Norm for better stability ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers) self.output_layer = nn.Linear(hidden_dim, output_dim) diff --git a/drug_discovery/mrna_therapeutics/__init__.py b/drug_discovery/mrna_therapeutics/__init__.py index 467ae47b51..859c3b6c36 100644 --- a/drug_discovery/mrna_therapeutics/__init__.py +++ b/drug_discovery/mrna_therapeutics/__init__.py @@ -3,4 +3,4 @@ Self-amplifying mRNA design. """ -from .mrna_optimizer import mRNAOptimizer, mRNAResult \ No newline at end of file +from .mrna_optimizer import mRNAOptimizer, mRNAResult diff --git a/drug_discovery/mrna_therapeutics/mrna_optimizer.py b/drug_discovery/mrna_therapeutics/mrna_optimizer.py index 8842c18211..dbdc49712f 100644 --- a/drug_discovery/mrna_therapeutics/mrna_optimizer.py +++ b/drug_discovery/mrna_therapeutics/mrna_optimizer.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any + @dataclass class mRNAResult: @@ -12,13 +12,13 @@ class mRNAResult: half_life: float = 0.0 success: bool = False + class mRNAOptimizer: - """mRNA sequence optimizer with saRNA (2026 breakthrough). - """ + """mRNA sequence optimizer with saRNA (2026 breakthrough).""" def optimize(self, antigen: str) -> mRNAResult: utr = "saRNA_UTR" expr = 95.0 immuno = 0.15 hl = 48.0 - return mRNAResult(antigen, utr, expr, immuno, hl, success=True) \ No newline at end of file + return mRNAResult(antigen, utr, expr, immuno, hl, success=True) diff --git a/drug_discovery/multi_omics/__init__.py b/drug_discovery/multi_omics/__init__.py index 8c826ca5f4..14ea2da107 100644 --- a/drug_discovery/multi_omics/__init__.py +++ b/drug_discovery/multi_omics/__init__.py @@ -40,16 +40,16 @@ ) __all__ = [ - "SingleCellLoader", - "SpatialTranscriptomicsLoader", + "ADMETConfig", + "ADMETPredictor", + "ADMETProfile", "CellData", - "HeterogeneousGraph", - "GraphNode", - "GraphEdge", "DrugTargetInteraction", - "NodeType", "EdgeType", - "ADMETPredictor", - "ADMETProfile", - "ADMETConfig", + "GraphEdge", + "GraphNode", + "HeterogeneousGraph", + "NodeType", + "SingleCellLoader", + "SpatialTranscriptomicsLoader", ] diff --git a/drug_discovery/multi_omics/admet_predictor.py b/drug_discovery/multi_omics/admet_predictor.py index 22ce503475..00f12b9fc4 100644 --- a/drug_discovery/multi_omics/admet_predictor.py +++ b/drug_discovery/multi_omics/admet_predictor.py @@ -22,7 +22,6 @@ import torch import torch.nn as nn import torch.nn.functional as F - from torch_geometric.data import Data from torch_geometric.nn import MessagePassing, global_max_pool, global_mean_pool @@ -282,7 +281,7 @@ def predict(self, smiles: str) -> ADMETProfile: except Exception as e: logger.error(f"ADMET prediction failed: {e}") - return ADMETProfile(clinical_risk_flags=[f"Prediction error: {str(e)}"]) + return ADMETProfile(clinical_risk_flags=[f"Prediction error: {e!s}"]) def _smiles_to_graph(self, mol) -> Data: """Convert RDKit molecule to PyTorch Geometric Data.""" diff --git a/drug_discovery/multi_omics/heterogeneous_graph.py b/drug_discovery/multi_omics/heterogeneous_graph.py index 9a7a222b77..c7519765d7 100644 --- a/drug_discovery/multi_omics/heterogeneous_graph.py +++ b/drug_discovery/multi_omics/heterogeneous_graph.py @@ -23,7 +23,6 @@ try: import torch import torch.nn as nn - from torch_geometric.data import HeteroData TORCH_GEOMETRIC_AVAILABLE = True diff --git a/drug_discovery/multi_omics/single_cell.py b/drug_discovery/multi_omics/single_cell.py index 8c24f2b32f..6dbd1d227d 100644 --- a/drug_discovery/multi_omics/single_cell.py +++ b/drug_discovery/multi_omics/single_cell.py @@ -11,6 +11,7 @@ from __future__ import annotations +import contextlib import logging from dataclasses import dataclass, field from typing import Any @@ -217,10 +218,7 @@ def get_cell_data(self, adata: Any) -> list[CellData]: for i, cell_id in enumerate(obs_names): # Gene expression expr = adata.x_data[i] if hasattr(adata.x_data, "__getitem__") else adata.x_data[i, :] - if hasattr(expr, "toarray"): - expr = expr.toarray().flatten() - else: - expr = np.array(expr).flatten() + expr = expr.toarray().flatten() if hasattr(expr, "toarray") else np.array(expr).flatten() # Cell type cell_type = "unknown" @@ -233,10 +231,8 @@ def get_cell_data(self, adata: Any) -> list[CellData]: metadata = {} for col in adata.obs.columns: if col not in ["cell_type", "celltype"]: - try: + with contextlib.suppress(Exception): metadata[col] = adata.obs[col].iloc[i] - except Exception: - pass cell = CellData( cell_id=str(cell_id), @@ -384,10 +380,7 @@ def _build_simple_spatial_graph( spatial_key: str, ) -> np.ndarray: """Build simple k-nearest-neighbor graph.""" - if spatial_key in adata.obsm: - coords = adata.obsm[spatial_key] - else: - coords = self._generate_spatial_coords(adata.n_obs) + coords = adata.obsm[spatial_key] if spatial_key in adata.obsm else self._generate_spatial_coords(adata.n_obs) n = coords.shape[0] distances = np.zeros((n, n)) diff --git a/drug_discovery/nanobotics/__init__.py b/drug_discovery/nanobotics/__init__.py index 1873ed3cde..c8417f99cf 100644 --- a/drug_discovery/nanobotics/__init__.py +++ b/drug_discovery/nanobotics/__init__.py @@ -2,4 +2,4 @@ from .swarm_logic import DNAGateSimulator, NanobotMARL -__all__ = ["NanobotMARL", "DNAGateSimulator"] +__all__ = ["DNAGateSimulator", "NanobotMARL"] diff --git a/drug_discovery/neuromorphic/__init__.py b/drug_discovery/neuromorphic/__init__.py index 3dd4a96b58..d572545ec3 100644 --- a/drug_discovery/neuromorphic/__init__.py +++ b/drug_discovery/neuromorphic/__init__.py @@ -3,4 +3,4 @@ from .compiler import SNNCompiler from .inference import NeuromorphicInferenceEngine -__all__ = ["SNNCompiler", "NeuromorphicInferenceEngine"] +__all__ = ["NeuromorphicInferenceEngine", "SNNCompiler"] diff --git a/drug_discovery/neuromorphic/compiler.py b/drug_discovery/neuromorphic/compiler.py index 8dcb6c5650..130162dfd7 100644 --- a/drug_discovery/neuromorphic/compiler.py +++ b/drug_discovery/neuromorphic/compiler.py @@ -48,7 +48,7 @@ def convert_to_spiking(self, ann_model: nn.Module, beta: float = 0.95) -> nn.Mod # activations with LIF layers. spiking_layers = [] - for name, module in ann_model.named_children(): + for _name, module in ann_model.named_children(): if isinstance(module, nn.Linear): spiking_layers.append(module) spiking_layers.append(snn.Leaky(beta=beta, spike_grad=surrogate.fast_sigmoid())) diff --git a/drug_discovery/optimization/__init__.py b/drug_discovery/optimization/__init__.py index fb1c27067c..9fa7f0fac9 100644 --- a/drug_discovery/optimization/__init__.py +++ b/drug_discovery/optimization/__init__.py @@ -23,8 +23,8 @@ "GaussianProcessSurrogate", "MOBOConfig", "MultiObjectiveBayesianOptimizer", - "is_pareto_efficient", "hypervolume_indicator", + "is_pareto_efficient", ] ) except ImportError: diff --git a/drug_discovery/physics/__init__.py b/drug_discovery/physics/__init__.py index 18a8f9534e..55dd654141 100644 --- a/drug_discovery/physics/__init__.py +++ b/drug_discovery/physics/__init__.py @@ -15,7 +15,7 @@ VinaBackend as VinaBackend, ) - __all__.extend(["DockingPipeline", "DockingConfig", "DockingResult", "VinaBackend"]) + __all__.extend(["DockingConfig", "DockingPipeline", "DockingResult", "VinaBackend"]) except ImportError: pass @@ -37,7 +37,7 @@ from drug_discovery.polyglot_integration import FEPResult as FEPResult from drug_discovery.polyglot_integration import PhysicsOracle as PhysicsOracle - __all__.extend(["PhysicsOracle", "FEPResult"]) + __all__.extend(["FEPResult", "PhysicsOracle"]) except ImportError: pass diff --git a/drug_discovery/physics/crystal_quality.py b/drug_discovery/physics/crystal_quality.py index 6045550f86..912839899d 100644 --- a/drug_discovery/physics/crystal_quality.py +++ b/drug_discovery/physics/crystal_quality.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any + @dataclass class CrystalResult: @@ -11,9 +11,10 @@ class CrystalResult: nmr_signal_boost: float = 0.0 # photo-CIDNP quality_grade: str = "high" + class CrystalEnhancer: """Crystal quality enhancement with photo-NMR (2025 FBDD breakthrough). - + Simulates hyperpolarization, XRPD/Raman polymorph analysis. """ @@ -24,4 +25,4 @@ def enhance_screen(self, fragments: list[str]) -> list[CrystalResult]: score = 0.95 boost = 40.0 # 30-50x results.append(CrystalResult(frag, rmsd, score, boost)) - return results \ No newline at end of file + return results diff --git a/drug_discovery/physics/docking.py b/drug_discovery/physics/docking.py index 34fbf5a0b1..2ffc8e1e9f 100644 --- a/drug_discovery/physics/docking.py +++ b/drug_discovery/physics/docking.py @@ -113,7 +113,7 @@ def dock_single(self, receptor, ligand, smiles=""): def dock_batch(self, receptor, ligands, smiles_list=None): sl = smiles_list or [""] * len(ligands) - return [self.dock_single(receptor, lig, s) for lig, s in zip(ligands, sl)] + return [self.dock_single(receptor, lig, s) for lig, s in zip(ligands, sl, strict=False)] @staticmethod def rank_results(results, top_k=10): diff --git a/drug_discovery/physics/md_simulator.py b/drug_discovery/physics/md_simulator.py index 8ed1f86141..1fa44011e4 100644 --- a/drug_discovery/physics/md_simulator.py +++ b/drug_discovery/physics/md_simulator.py @@ -244,19 +244,19 @@ def calculate_energy(self, smiles: str) -> float | None: return None mol = Chem.AddHs(mol) - embed_molecule = getattr(AllChem, "EmbedMolecule") + embed_molecule = AllChem.EmbedMolecule success = embed_molecule(mol, randomSeed=42) if success == -1: return None if self.method.lower() == "mmff94": - mmff_props = getattr(AllChem, "MMFFGetMoleculeProperties") - mmff_force_field = getattr(AllChem, "MMFFGetMoleculeForceField") + mmff_props = AllChem.MMFFGetMoleculeProperties + mmff_force_field = AllChem.MMFFGetMoleculeForceField props = mmff_props(mol) ff = mmff_force_field(mol, props) energy = ff.CalcEnergy() elif self.method.lower() == "uff": - uff_force_field = getattr(AllChem, "UFFGetMoleculeForceField") + uff_force_field = AllChem.UFFGetMoleculeForceField ff = uff_force_field(mol) energy = ff.CalcEnergy() else: @@ -288,20 +288,20 @@ def optimize_geometry(self, smiles: str, max_iters: int = 200) -> tuple[str | No return None, None mol = Chem.AddHs(mol) - embed_molecule = getattr(AllChem, "EmbedMolecule") + embed_molecule = AllChem.EmbedMolecule success = embed_molecule(mol, randomSeed=42) if success == -1: return None, None if self.method.lower() == "mmff94": - mmff_props = getattr(AllChem, "MMFFGetMoleculeProperties") - mmff_force_field = getattr(AllChem, "MMFFGetMoleculeForceField") + mmff_props = AllChem.MMFFGetMoleculeProperties + mmff_force_field = AllChem.MMFFGetMoleculeForceField props = mmff_props(mol) ff = mmff_force_field(mol, props) ff.Minimize(maxIts=max_iters) energy = ff.CalcEnergy() elif self.method.lower() == "uff": - uff_force_field = getattr(AllChem, "UFFGetMoleculeForceField") + uff_force_field = AllChem.UFFGetMoleculeForceField ff = uff_force_field(mol) ff.Minimize(maxIts=max_iters) energy = ff.CalcEnergy() diff --git a/drug_discovery/pipeline.py b/drug_discovery/pipeline.py index fd4f1c7e7c..d43f59d0e0 100644 --- a/drug_discovery/pipeline.py +++ b/drug_discovery/pipeline.py @@ -7,15 +7,16 @@ import sys from collections.abc import Sequence from pathlib import Path -from typing import Any, Dict, List, Optional, cast +from typing import Any, cast import numpy as np import pandas as pd import torch from torch.utils.data import DataLoader - from torch_geometric.loader import DataLoader as GeometricDataLoader +from .biomarker_discovery import BiomarkerMLDiscovery +from .causal_discovery import CausalGraph, CausalInference from .data import ( DataCollector, MolecularDataset, @@ -24,14 +25,10 @@ train_test_split_molecular, ) from .evaluation import ADMETPredictor, ModelEvaluator, PropertyPredictor, TorchDrugScorer -from .models import EnsembleModel, MolecularGNN, MolecularTransformer, MolecularGIN, ModernMolecularTransformer +from .models import EnsembleModel, ModernMolecularTransformer, MolecularGIN, MolecularGNN, MolecularTransformer from .physics import DiffDockAdapter, OpenFoldAdapter, OpenMMAdapter from .synthesis import MolecularTransformerAdapter, PistachioDatasets from .training import SelfLearningTrainer -from .biomarker_discovery import BiomarkerMLDiscovery, BiomarkerStatisticalAnalysis -from .explainability import GraphExplainer, FingerprintExplainer -from .causal_discovery import CausalGraph, CausalInference -from .lead_optimization import LeadMCTSOptimizer, LeadRLOptimizer class DrugDiscoveryPipeline: @@ -72,7 +69,7 @@ def __init__( self.trainer = None self.property_predictor = None self.learnable_docking = None - + # Advanced Modules self.explainers = {} self.causal_graph = CausalGraph() @@ -173,10 +170,7 @@ def prepare_datasets( print("\n=== Data Preparation Phase ===") # Determine featurization based on model type - if self.model_type in ["gnn", "gin"]: - featurization = "graph" - else: - featurization = "fingerprint" + featurization = "graph" if self.model_type in ["gnn", "gin"] else "fingerprint" # Create dataset dataset = MolecularDataset(data=data, smiles_col=smiles_col, target_col=target_col, featurization=featurization) @@ -838,60 +832,58 @@ def run_precision_medicine_workflow(self, patient_data: pd.DataFrame, target_col 3. Match drugs to cluster profiles. """ print("\n=== Running Precision Medicine Workflow ===") - + from .precision_medicine import PatientStratifier - from .biomarker_discovery import BiomarkerMLDiscovery - + # 1. Stratification stratifier = PatientStratifier(patient_data) features = [col for col in patient_data.columns if col != target_col] clusters = stratifier.stratify_patients(features) cluster_info = stratifier.get_cluster_characteristics(clusters) - + # 2. Biomarker Discovery for each cluster discovery = BiomarkerMLDiscovery(patient_data, target_col) biomarkers = discovery.rank_features_by_importance(features) - + print(f"Detected {len(cluster_info)} patient clusters.") return { "clusters": clusters.to_dict(), "cluster_characteristics": cluster_info.to_dict(), - "top_biomarkers": biomarkers.head(10).to_dict() + "top_biomarkers": biomarkers.head(10).to_dict(), } - def refine_lead_candidates(self, candidates_smiles: List[str], target_protein_pdb: Optional[str] = None) -> List[Dict[str, Any]]: + def refine_lead_candidates( + self, candidates_smiles: list[str], target_protein_pdb: str | None = None + ) -> list[dict[str, Any]]: """ Refine a list of candidates using the Ensemble Refiner. """ from .lead_optimization import EnsembleRefiner - + if self.property_predictor is None: self.property_predictor = PropertyPredictor(self.model, self.device) - + refiner = EnsembleRefiner( property_predictor=self.property_predictor, admet_predictor=self.admet_predictor, - physics_simulator=self # Pipeline itself can act as adapter if it has OpenMM calls + physics_simulator=self, # Pipeline itself can act as adapter if it has OpenMM calls ) - + return refiner.rank_candidates(candidates_smiles, target_protein_pdb) def perform_causal_reasoning(self, data: pd.DataFrame, treatment: str, outcome: str) -> dict[str, Any]: """ Identify causal effects in a dataset. """ - from .causal_discovery import CausalInference, CausalGraph - + from .causal_discovery import CausalGraph + # 1. Structure discovery graph = CausalGraph() graph.discover_from_data(data) - + # 2. Inference inference = CausalInference(data) confounders = [col for col in data.columns if col not in [treatment, outcome]] ate = inference.estimate_treatment_effect(treatment, outcome, confounders) - - return { - "average_treatment_effect": ate, - "causal_graph_summary": graph.graph.number_of_edges() - } + + return {"average_treatment_effect": ate, "causal_graph_summary": graph.graph.number_of_edges()} diff --git a/drug_discovery/pipeline/__init__.py b/drug_discovery/pipeline/__init__.py index aad444b234..6bf3a42b1d 100644 --- a/drug_discovery/pipeline/__init__.py +++ b/drug_discovery/pipeline/__init__.py @@ -32,11 +32,11 @@ DrugDiscoveryPipeline = None __all__ = [ - "DrugDiscoveryPipeline", - "StreamingDataPipeline", - "PipelineOrchestrator", + "DataBatch", "DataQualityMonitor", + "DrugDiscoveryPipeline", "FaultTolerantExecutor", - "DataBatch", "PipelineCheckpoint", + "PipelineOrchestrator", + "StreamingDataPipeline", ] diff --git a/drug_discovery/pipeline/autonomous_pipeline.py b/drug_discovery/pipeline/autonomous_pipeline.py index 1c9959a321..9f5660eb4c 100644 --- a/drug_discovery/pipeline/autonomous_pipeline.py +++ b/drug_discovery/pipeline/autonomous_pipeline.py @@ -154,10 +154,7 @@ async def execute_with_retry( if attempt < self.max_retries - 1: # Calculate delay - if self.exponential_backoff: - delay = self.retry_delay * (2**attempt) - else: - delay = self.retry_delay + delay = self.retry_delay * 2**attempt if self.exponential_backoff else self.retry_delay logger.info(f"Retrying in {delay} seconds...") await asyncio.sleep(delay) @@ -299,10 +296,11 @@ async def stream_data_batches( process_function, batch_df, ) - + # Elite filtering (20000x better candidates) if elite_filter and "smiles" in processed_batch.columns: from drug_discovery.safety import SmilesValidator + validator = SmilesValidator() mask = processed_batch["smiles"].apply(validator.is_elite_smiles) processed_batch = processed_batch[mask].reset_index(drop=True) @@ -489,7 +487,7 @@ async def run_pipeline(p): # Run all pipelines concurrently stats_list = await asyncio.gather(*tasks, return_exceptions=True) - for pipeline_id, stats in zip(self.pipelines.keys(), stats_list): + for pipeline_id, stats in zip(self.pipelines.keys(), stats_list, strict=False): if isinstance(stats, Exception): logger.error(f"Pipeline {pipeline_id} failed: {stats}") else: diff --git a/drug_discovery/polyglot_integration.py b/drug_discovery/polyglot_integration.py index 32f3a55e9b..3b26edd604 100644 --- a/drug_discovery/polyglot_integration.py +++ b/drug_discovery/polyglot_integration.py @@ -135,11 +135,14 @@ def _fep_openmm( energies: list[float] = [] for lam in lambda_values: - ctx = mm.Context(system, mm.LangevinMiddleIntegrator( - temperature * unit.kelvin, - 1.0 / unit.picosecond, - timestep * unit.femtoseconds, - )) + ctx = mm.Context( + system, + mm.LangevinMiddleIntegrator( + temperature * unit.kelvin, + 1.0 / unit.picosecond, + timestep * unit.femtoseconds, + ), + ) ctx.setPositions([mm.Vec3(0.0, 0.0, 0.0)] * unit.nanometers) ctx.getIntegrator().step(min(steps_per_window, 100)) state = ctx.getState(getEnergy=True) @@ -321,7 +324,7 @@ async def score_batch(self, smiles_list: Sequence[str]) -> list[FEPResult]: if to_compute: uncached_smiles = [smi for _, smi in to_compute] computed = await self._compute_with_retry(uncached_smiles) - for (idx, smi), result in zip(to_compute, computed): + for (idx, smi), result in zip(to_compute, computed, strict=False): results[idx] = result if self._cache is not None and result.success: self._cache[self._cache_key(smi)] = result @@ -359,7 +362,7 @@ async def _compute_with_retry(self, smiles_list: list[str]) -> list[FEPResult]: ) retry_smiles = [smiles_list[i] for i in failed_indices] retry_results = await self._dispatch(retry_smiles) - for fi, rr in zip(failed_indices, retry_results): + for fi, rr in zip(failed_indices, retry_results, strict=False): if rr.success: results[fi] = rr @@ -391,9 +394,7 @@ async def _score_batch_ray(self, smiles_list: Sequence[str]) -> list[FEPResult]: for smi in smiles_list ] - raw_results = await asyncio.gather( - *[asyncio.wrap_future(f.future()) for f in futures] - ) + raw_results = await asyncio.gather(*[asyncio.wrap_future(f.future()) for f in futures]) return [self._dict_to_result(d) for d in raw_results] # ------------------------------------------------------------------ diff --git a/drug_discovery/precision_medicine/genomic_matcher.py b/drug_discovery/precision_medicine/genomic_matcher.py index 4980716c7e..47848359f7 100644 --- a/drug_discovery/precision_medicine/genomic_matcher.py +++ b/drug_discovery/precision_medicine/genomic_matcher.py @@ -1,18 +1,16 @@ -from typing import List, Dict, Any - class GenomicDrugMatcher: """ Match drugs to specific genomic variants. """ - - def __init__(self, variant_drug_database: Dict[str, List[str]]): + + def __init__(self, variant_drug_database: dict[str, list[str]]): """ Args: variant_drug_database: Mapping from variant (e.g., 'EGFR T790M') to list of effective drugs. """ self.db = variant_drug_database - - def match_drugs_to_patient(self, patient_variants: List[str]) -> List[str]: + + def match_drugs_to_patient(self, patient_variants: list[str]) -> list[str]: """ Find drugs that match the patient's genetic profile. """ diff --git a/drug_discovery/precision_medicine/patient_stratification.py b/drug_discovery/precision_medicine/patient_stratification.py index 29a83c622a..60e4e847b2 100644 --- a/drug_discovery/precision_medicine/patient_stratification.py +++ b/drug_discovery/precision_medicine/patient_stratification.py @@ -1,16 +1,16 @@ import pandas as pd from sklearn.cluster import KMeans -from typing import List, Dict, Any + class PatientStratifier: """ Cluster patients based on multi-omics or clinical data. """ - + def __init__(self, data: pd.DataFrame): self.data = data - - def stratify_patients(self, features: List[str], num_clusters: int = 3) -> pd.Series: + + def stratify_patients(self, features: list[str], num_clusters: int = 3) -> pd.Series: """ Group patients into clusters. """ @@ -24,5 +24,5 @@ def get_cluster_characteristics(self, clusters: pd.Series) -> pd.DataFrame: Analyze the mean values of features for each cluster. """ df_with_clusters = self.data.copy() - df_with_clusters['cluster'] = clusters - return df_with_clusters.groupby('cluster').mean() + df_with_clusters["cluster"] = clusters + return df_with_clusters.groupby("cluster").mean() diff --git a/drug_discovery/qml/__init__.py b/drug_discovery/qml/__init__.py index 5eba29f16b..a791280abe 100644 --- a/drug_discovery/qml/__init__.py +++ b/drug_discovery/qml/__init__.py @@ -36,21 +36,21 @@ ) __all__ = [ + "AWSBraketDriver", # Active Space "ActiveSpaceApproximator", "ActiveSpaceResult", + "ErrorMitigationConfig", + "HardwareEfficientAnsatz", + "LocalSimulator", "MolecularOrbitals", + # Drivers + "QuantumDriver", + "QuantumSimulator", # VQE "VQECircuit", "VQEResult", - "HardwareEfficientAnsatz", + "ZNEResult", # Error Mitigation "ZeroNoiseExtrapolation", - "ZNEResult", - "ErrorMitigationConfig", - # Drivers - "QuantumDriver", - "QuantumSimulator", - "AWSBraketDriver", - "LocalSimulator", ] diff --git a/drug_discovery/qml/quantum_chemistry.py b/drug_discovery/qml/quantum_chemistry.py index e5ee98c732..4d1221b2d4 100644 --- a/drug_discovery/qml/quantum_chemistry.py +++ b/drug_discovery/qml/quantum_chemistry.py @@ -264,7 +264,7 @@ def approximate_active_space( n_active_electrons = (n_active_electrons // 2) * 2 # Extract active space - active_energies, active_coeffs, active_occ, core_energy = orbitals.get_active_space( + active_energies, active_coeffs, _active_occ, core_energy = orbitals.get_active_space( n_active_electrons, n_active_orbitals ) diff --git a/drug_discovery/qml/quantum_driver.py b/drug_discovery/qml/quantum_driver.py index e0aa4ca563..1aac4a7e92 100644 --- a/drug_discovery/qml/quantum_driver.py +++ b/drug_discovery/qml/quantum_driver.py @@ -13,6 +13,7 @@ from __future__ import annotations +import contextlib import logging from abc import ABC, abstractmethod from dataclasses import dataclass, field @@ -380,10 +381,7 @@ def measure(self, observable: str, n_shots: int = 1000) -> QuantumResult: start = time.time() try: - if self.local: - device = self._LocalSimulator() - else: - device = self._AwsDevice(self.device_arn) + device = self._LocalSimulator() if self.local else self._AwsDevice(self.device_arn) task = device.run(self._circuit, shots=n_shots) result = task.result() @@ -525,10 +523,8 @@ def available_backends(self) -> list[str]: """List available quantum backends.""" backends = ["local"] - try: + with contextlib.suppress(ImportError): backends.append("aws_braket") - except ImportError: - pass return backends diff --git a/drug_discovery/qml/vqe.py b/drug_discovery/qml/vqe.py index b3a8d46586..b29e6610ad 100644 --- a/drug_discovery/qml/vqe.py +++ b/drug_discovery/qml/vqe.py @@ -144,7 +144,7 @@ def circuit_template(params): # Entangling layer + subsequent rotations param_idx = self.n_qubits - for layer in range(self.n_layers): + for _layer in range(self.n_layers): # Entangling gates if self.entanglement == "linear": for i in range(self.n_qubits - 1): @@ -179,7 +179,7 @@ def _simulate_circuit(self, parameters: np.ndarray) -> float: total += np.cos(theta) # Simplified RY effect param_idx += 1 - for layer in range(self.n_layers): + for _layer in range(self.n_layers): # Simplified entanglement effect for i in range(self.n_qubits - 1): total += 0.1 * np.sin(parameters[param_idx + i]) @@ -302,7 +302,7 @@ def circuit(params): self._apply_ansatz(params) # Measure expectation results = [] - for term in hamiltonian.keys(): + for _term in hamiltonian: # Simplified measurement results.append(qml.expval(qml.PauliZ(0))) return sum(results) / len(results) if results else 0.0 @@ -319,7 +319,7 @@ def _apply_ansatz(self, params: torch.Tensor) -> None: qml.RY(params[param_idx], wires=i) param_idx += 1 - for layer in range(self.n_layers): + for _layer in range(self.n_layers): # Entangling gates for i in range(self.n_qubits - 1): qml.CNOT(wires=[i, i + 1]) diff --git a/drug_discovery/rfdiffusion/__init__.py b/drug_discovery/rfdiffusion/__init__.py index eedbe27a8a..19b9470a2a 100644 --- a/drug_discovery/rfdiffusion/__init__.py +++ b/drug_discovery/rfdiffusion/__init__.py @@ -3,6 +3,6 @@ SE(3) diffusion for motif scaffolding. """ -from .protein_design import RFDiffusionDesigner, RFDesignResult +from .protein_design import RFDesignResult, RFDiffusionDesigner -__all__ = ["RFDiffusionDesigner", "RFDesignResult"] \ No newline at end of file +__all__ = ["RFDesignResult", "RFDiffusionDesigner"] diff --git a/drug_discovery/rfdiffusion/protein_design.py b/drug_discovery/rfdiffusion/protein_design.py index 3761581ceb..a515342efe 100644 --- a/drug_discovery/rfdiffusion/protein_design.py +++ b/drug_discovery/rfdiffusion/protein_design.py @@ -1,7 +1,8 @@ from __future__ import annotations +from collections.abc import Sequence from dataclasses import dataclass -from typing import Any, Sequence + @dataclass class RFDesignResult: @@ -13,9 +14,10 @@ class RFDesignResult: success: bool = False error: str | None = None + class RFDiffusionDesigner: """Proxy for RFdiffusion protein design using diffusion models. - + Falls back to simple helix insertion if deps unavailable. Supports Ray batch design. """ @@ -29,9 +31,9 @@ def design_batch(self, motifs: Sequence[str]) -> list[RFDesignResult]: for motif in motifs: try: # Mock design: append helix - design = motif + "HELI" * (len(motif)//3 + 1) + design = motif + "HELI" * (len(motif) // 3 + 1) rmsd = 1.5 + len(motif) % 5 * 0.5 # mock results.append(RFDesignResult(motif, design, rmsd_recovery=rmsd, confidence=0.85, success=True)) except Exception as e: results.append(RFDesignResult(motif, "", error=str(e))) - return results \ No newline at end of file + return results diff --git a/drug_discovery/safety/end_to_end_pipeline.py b/drug_discovery/safety/end_to_end_pipeline.py index 53cf6bc098..7426797442 100644 --- a/drug_discovery/safety/end_to_end_pipeline.py +++ b/drug_discovery/safety/end_to_end_pipeline.py @@ -122,10 +122,7 @@ def run( # ---------------------------------------------------------- # Step 1: Generate # ---------------------------------------------------------- - if seed_smiles: - raw_smiles = list(seed_smiles) - else: - raw_smiles = self._generate(n) + raw_smiles = list(seed_smiles) if seed_smiles else self._generate(n) result.candidates_generated = len(raw_smiles) logger.info("Step 1: Generated %d candidates", len(raw_smiles)) @@ -194,14 +191,17 @@ def _generate(self, n: int) -> list[str]: # Default: diverse drug-like seed pool pool = [ - "CCO", "c1ccccc1", "CC(=O)O", + "CCO", + "c1ccccc1", + "CC(=O)O", "CC(=O)Oc1ccccc1C(=O)O", "CC(C)Cc1ccc(cc1)C(C)C(=O)O", "CN1C=NC2=C1C(=O)N(C(=O)N2C)C", "CC12CCC3C(C1CCC2O)CCC4=CC(=O)CCC34C", "OC(=O)c1ccccc1O", "CC(=O)NC1=CC=C(C=C1)O", - "C1CCCCC1", "c1ccncc1", + "C1CCCCC1", + "c1ccncc1", "C1=CC=C(C=C1)C(=O)O", "OC1=CC=CC=C1", "c1ccc(cc1)N", @@ -233,10 +233,7 @@ def _surrogate_filter(self, smiles_list: list[str]) -> list[str]: def _oracle_score(self, smiles_list: list[str]) -> list[dict[str, Any]]: if self._physics_oracle is not None: results = self._physics_oracle.score_batch_sync(smiles_list) - return [ - {"smiles": r.smiles, "delta_g": r.delta_g or 0.0} - for r in results - ] + return [{"smiles": r.smiles, "delta_g": r.delta_g or 0.0} for r in results] # Mock scoring import hashlib diff --git a/drug_discovery/safety/environmental_tests.py b/drug_discovery/safety/environmental_tests.py index 23470d9b45..5ffff9833f 100644 --- a/drug_discovery/safety/environmental_tests.py +++ b/drug_discovery/safety/environmental_tests.py @@ -3,13 +3,14 @@ These are lightweight heuristic/estimator modules useful for gating compounds before heavy experimental assays. Intended to be used as part of QA pipelines. """ + from __future__ import annotations -from typing import Dict, Any import math +from typing import Any -def estimate_ph_stability(smiles: str, ph: float) -> Dict[str, Any]: +def estimate_ph_stability(smiles: str, ph: float) -> dict[str, Any]: """Estimate percent remaining after 24h at given pH using heuristics. This is a heuristic fallback when experimental data is not available. @@ -17,10 +18,10 @@ def estimate_ph_stability(smiles: str, ph: float) -> Dict[str, Any]: """ # Heuristic: acids degrade in high pH, bases degrade in low pH. score = 0.8 - if any(x in smiles for x in ['C(=O)O', 'CO2H', 'O=C(O)']): # carboxylic acid moieties + if any(x in smiles for x in ["C(=O)O", "CO2H", "O=C(O)"]): # carboxylic acid moieties # acids less stable at high pH score -= max(0.0, (ph - 7.0) * 0.08) - if any(x in smiles for x in ['N', 'NH', 'N(']): + if any(x in smiles for x in ["N", "NH", "N("]): # basic amines less stable at low pH score -= max(0.0, (7.0 - ph) * 0.06) # clamp @@ -28,22 +29,22 @@ def estimate_ph_stability(smiles: str, ph: float) -> Dict[str, Any]: return {"percent_remaining": round(score * 100.0, 1), "confidence": 0.45} -def estimate_plasma_binding(smiles: str) -> Dict[str, Any]: +def estimate_plasma_binding(smiles: str) -> dict[str, Any]: """Estimate fraction bound to plasma proteins (heuristic). Returns {'fraction_bound': 0-1, 'confidence': 0-1} """ # Heuristic: high logP and aromatic rings increase plasma binding - aromatic = smiles.count('c') + smiles.count('C') // 4 - logp_est = 1.0 + aromatic * 0.4 - smiles.count('N') * 0.3 + aromatic = smiles.count("c") + smiles.count("C") // 4 + logp_est = 1.0 + aromatic * 0.4 - smiles.count("N") * 0.3 # map to 0-1 - frac = 1.0 - 1.0 / (1.0 + math.exp((logp_est - 2.5))) + frac = 1.0 - 1.0 / (1.0 + math.exp(logp_est - 2.5)) frac = max(0.0, min(0.99, frac)) confidence = 0.4 return {"fraction_bound": round(frac, 3), "confidence": confidence} -def run_environmental_tests(smiles: str, ph_values: tuple[float, ...] = (1.2, 4.5, 7.4, 9.0)) -> Dict[str, Any]: +def run_environmental_tests(smiles: str, ph_values: tuple[float, ...] = (1.2, 4.5, 7.4, 9.0)) -> dict[str, Any]: results = {"smiles": smiles, "ph_profiles": {}, "plasma_binding": None} for ph in ph_values: results["ph_profiles"][str(ph)] = estimate_ph_stability(smiles, ph) diff --git a/drug_discovery/safety/parametrized_toxicity_gate.py b/drug_discovery/safety/parametrized_toxicity_gate.py index 7cd423cd84..5284fbb694 100644 --- a/drug_discovery/safety/parametrized_toxicity_gate.py +++ b/drug_discovery/safety/parametrized_toxicity_gate.py @@ -6,30 +6,30 @@ from __future__ import annotations -from dataclasses import dataclass, field -from typing import Any, Optional +from dataclasses import dataclass +from typing import Any @dataclass class ToxicityThresholdConfig: """Parametrized toxicity thresholds (no hardcoding).""" - + # hERG-related thresholds herg_threshold: float = 0.3 # hERG inhibition probability herg_warning_threshold: float = 0.2 # Trigger additional scrutiny cyp3a4_inhibition_threshold: float = 0.5 # CYP3A4 substrate probability cyp2d6_inhibition_threshold: float = 0.4 # CYP2D6 inhibition - + # Mutagenicity thresholds ames_threshold: float = 0.3 # Ames mutagenicity - + # Hepatotoxicity thresholds hepatotox_threshold: float = 0.4 # Hepatotoxicity probable drug_induced_liver_injury_threshold: float = 0.35 # DILI risk - + # Cytotoxicity thresholds cytotox_threshold: float = 0.4 # General cytotoxicity - + # Physicochemical thresholds logp_max: float = 5.0 logp_min: float = -1.0 @@ -38,19 +38,19 @@ class ToxicityThresholdConfig: mw_max: float = 500.0 mw_min: float = 50.0 rotatable_bonds_max: int = 10 - + # Bioavailability flags require_lipinski_compliance: bool = True lipinski_violations_max: int = 1 # Allow up to 1 violation - + # Regulatory flags require_no_known_toxicophores: bool = True allow_experimental_optimization: bool = False - + # Quality flags min_prediction_confidence: float = 0.5 # 0-1 scale require_vendor_approval: bool = False - + def __post_init__(self): """Validate threshold ranges.""" if not (0 <= self.herg_threshold <= 1): @@ -63,11 +63,11 @@ def __post_init__(self): raise ValueError("logp_min must be < logp_max") if self.mw_min >= self.mw_max: raise ValueError("mw_min must be < mw_max") - + @classmethod def from_regulatory_tier(cls, tier: str) -> ToxicityThresholdConfig: """Create thresholds for regulatory submission tier. - + Args: tier: One of ['discovery', 'lead_optimization', 'ind', 'nda'] """ @@ -105,23 +105,23 @@ def from_regulatory_tier(cls, tier: str) -> ToxicityThresholdConfig: class ParametrizedToxicityGate: """Toxicity gate with configurable thresholds. - + Usage:: - + # Standard thresholds gate = ParametrizedToxicityGate() - + # IND submission thresholds (stricter) config = ToxicityThresholdConfig.from_regulatory_tier("ind") gate = ParametrizedToxicityGate(config) - + # Custom thresholds config = ToxicityThresholdConfig( herg_threshold=0.2, ames_threshold=0.1, ) gate = ParametrizedToxicityGate(config) - + result = gate.evaluate( herg_prob=0.1, ames_prob=0.05, @@ -130,26 +130,26 @@ class ParametrizedToxicityGate: if not result['passed']: print(f"Rejected: {result['reasons']}") """ - - def __init__(self, config: Optional[ToxicityThresholdConfig] = None): + + def __init__(self, config: ToxicityThresholdConfig | None = None): """Initialize with thresholds (no hardcoding).""" self.config = config or ToxicityThresholdConfig() - + def evaluate( self, - herg_prob: Optional[float] = None, - ames_prob: Optional[float] = None, - hepatotox_prob: Optional[float] = None, - cytotox_prob: Optional[float] = None, - logp: Optional[float] = None, - tpsa: Optional[float] = None, - mw: Optional[float] = None, - rotatable_bonds: Optional[int] = None, - confidence: Optional[float] = None, + herg_prob: float | None = None, + ames_prob: float | None = None, + hepatotox_prob: float | None = None, + cytotox_prob: float | None = None, + logp: float | None = None, + tpsa: float | None = None, + mw: float | None = None, + rotatable_bonds: int | None = None, + confidence: float | None = None, **kwargs: Any, ) -> dict[str, Any]: """Evaluate molecule against configured thresholds. - + Args: herg_prob: hERG inhibition probability [0-1] ames_prob: Ames mutagenicity probability [0-1] @@ -161,97 +161,79 @@ def evaluate( rotatable_bonds: Number of rotatable bonds confidence: Model prediction confidence [0-1] **kwargs: Additional properties - + Returns: Dictionary with 'passed' (bool) and 'reasons' (list of rejection reasons) """ passed = True reasons = [] warnings = [] - + # Check confidence - if confidence is not None: - if confidence < self.config.min_prediction_confidence: - passed = False - reasons.append( - f"Prediction confidence too low: {confidence:.2f} " - f"(minimum: {self.config.min_prediction_confidence})" - ) - + if confidence is not None and confidence < self.config.min_prediction_confidence: + passed = False + reasons.append( + f"Prediction confidence too low: {confidence:.2f} " + f"(minimum: {self.config.min_prediction_confidence})" + ) + # Check hERG if herg_prob is not None: if herg_prob > self.config.herg_threshold: passed = False reasons.append( - f"hERG inhibition too high: {herg_prob:.3f} " - f"(threshold: {self.config.herg_threshold})" + f"hERG inhibition too high: {herg_prob:.3f} " f"(threshold: {self.config.herg_threshold})" ) elif herg_prob > self.config.herg_warning_threshold: warnings.append( f"hERG inhibition warning: {herg_prob:.3f} " f"(caution threshold: {self.config.herg_warning_threshold})" ) - + # Check Ames - if ames_prob is not None: - if ames_prob > self.config.ames_threshold: - passed = False - reasons.append( - f"Ames mutagenicity too high: {ames_prob:.3f} " - f"(threshold: {self.config.ames_threshold})" - ) - + if ames_prob is not None and ames_prob > self.config.ames_threshold: + passed = False + reasons.append(f"Ames mutagenicity too high: {ames_prob:.3f} " f"(threshold: {self.config.ames_threshold})") + # Check hepatotoxicity - if hepatotox_prob is not None: - if hepatotox_prob > self.config.hepatotox_threshold: - passed = False - reasons.append( - f"Hepatotoxicity too high: {hepatotox_prob:.3f} " - f"(threshold: {self.config.hepatotox_threshold})" - ) - + if hepatotox_prob is not None and hepatotox_prob > self.config.hepatotox_threshold: + passed = False + reasons.append( + f"Hepatotoxicity too high: {hepatotox_prob:.3f} " f"(threshold: {self.config.hepatotox_threshold})" + ) + # Check cytotoxicity - if cytotox_prob is not None: - if cytotox_prob > self.config.cytotox_threshold: - passed = False - reasons.append( - f"Cytotoxicity too high: {cytotox_prob:.3f} " - f"(threshold: {self.config.cytotox_threshold})" - ) - + if cytotox_prob is not None and cytotox_prob > self.config.cytotox_threshold: + passed = False + reasons.append( + f"Cytotoxicity too high: {cytotox_prob:.3f} " f"(threshold: {self.config.cytotox_threshold})" + ) + # Check physicochemical properties - if logp is not None: - if not (self.config.logp_min <= logp <= self.config.logp_max): - passed = False - reasons.append( - f"LogP out of range: {logp:.2f} " - f"(allowed: [{self.config.logp_min}, {self.config.logp_max}])" - ) - - if tpsa is not None: - if not (self.config.tpsa_min <= tpsa <= self.config.tpsa_max): - passed = False - reasons.append( - f"TPSA out of range: {tpsa:.1f} " - f"(allowed: [{self.config.tpsa_min}, {self.config.tpsa_max}])" - ) - - if mw is not None: - if not (self.config.mw_min <= mw <= self.config.mw_max): - passed = False - reasons.append( - f"Molecular weight out of range: {mw:.1f} " - f"(allowed: [{self.config.mw_min}, {self.config.mw_max}])" - ) - - if rotatable_bonds is not None: - if rotatable_bonds > self.config.rotatable_bonds_max: - passed = False - reasons.append( - f"Too many rotatable bonds: {rotatable_bonds} " - f"(maximum: {self.config.rotatable_bonds_max})" - ) - + if logp is not None and not (self.config.logp_min <= logp <= self.config.logp_max): + passed = False + reasons.append( + f"LogP out of range: {logp:.2f} " f"(allowed: [{self.config.logp_min}, {self.config.logp_max}])" + ) + + if tpsa is not None and not (self.config.tpsa_min <= tpsa <= self.config.tpsa_max): + passed = False + reasons.append( + f"TPSA out of range: {tpsa:.1f} " f"(allowed: [{self.config.tpsa_min}, {self.config.tpsa_max}])" + ) + + if mw is not None and not (self.config.mw_min <= mw <= self.config.mw_max): + passed = False + reasons.append( + f"Molecular weight out of range: {mw:.1f} " f"(allowed: [{self.config.mw_min}, {self.config.mw_max}])" + ) + + if rotatable_bonds is not None and rotatable_bonds > self.config.rotatable_bonds_max: + passed = False + reasons.append( + f"Too many rotatable bonds: {rotatable_bonds} " f"(maximum: {self.config.rotatable_bonds_max})" + ) + return { "passed": passed, "reasons": reasons, @@ -262,10 +244,10 @@ def evaluate( "hepatotox": self.config.hepatotox_threshold, }, } - + def update_thresholds(self, **kwargs: Any) -> None: """Update threshold values without recreating object. - + Args: **kwargs: Threshold parameters to update """ diff --git a/drug_discovery/safety/smiles_validator.py b/drug_discovery/safety/smiles_validator.py index 82fe9e937a..796ad8afaa 100644 --- a/drug_discovery/safety/smiles_validator.py +++ b/drug_discovery/safety/smiles_validator.py @@ -243,7 +243,7 @@ def is_elite_smiles(self, smiles: str) -> bool: res = self.validate(smiles) if not res.passed: return False - + if not _RDKIT: return res.is_valid @@ -259,19 +259,15 @@ def is_elite_smiles(self, smiles: str) -> bool: for atom in mol.GetAtoms(): if atom.GetSymbol() == "C" and atom.GetTotalValence() > 4: return False - + # Avoid extremely reactive groups reactive_patterns = [ - "[F,Cl,Br,I][F,Cl,Br,I]", # Halogen-halogen - "[O,N]-[O,N]", # Peroxides/Hydrazines (sometimes okay but often unstable) - "C#C-C#C", # Polyynes - "[S,C]=[O,S]=O", # Ketenes/Sulfines + "[F,Cl,Br,I][F,Cl,Br,I]", # Halogen-halogen + "[O,N]-[O,N]", # Peroxides/Hydrazines (sometimes okay but often unstable) + "C#C-C#C", # Polyynes + "[S,C]=[O,S]=O", # Ketenes/Sulfines ] - for pattern in reactive_patterns: - if mol.HasSubstructMatch(Chem.MolFromSmarts(pattern)): - return False - - return True + return all(not mol.HasSubstructMatch(Chem.MolFromSmarts(pattern)) for pattern in reactive_patterns) # ------------------------------------------------------------------ # Heuristic fallback (no RDKit) diff --git a/drug_discovery/safety/strict_compliance_gate.py b/drug_discovery/safety/strict_compliance_gate.py index fe8040d315..ffcdce3a28 100644 --- a/drug_discovery/safety/strict_compliance_gate.py +++ b/drug_discovery/safety/strict_compliance_gate.py @@ -15,10 +15,11 @@ import hashlib import json import logging +from collections.abc import Sequence from dataclasses import dataclass, field from datetime import datetime from enum import Enum -from typing import Any, Optional, Sequence +from typing import Any logger = logging.getLogger(__name__) @@ -72,9 +73,9 @@ class ComplianceCheckResult: compliance_level: ComplianceLevel severity: str # "critical", "major", "minor", "info" message: str - value: Optional[float] = None - threshold: Optional[float] = None - remediation: Optional[str] = None + value: float | None = None + threshold: float | None = None + remediation: str | None = None timestamp: datetime = field(default_factory=datetime.utcnow) @@ -96,7 +97,7 @@ def verify_integrity(self, expected_hash: str) -> bool: self.integrity_verified = True return True self.tampering_detected = True - logger.warning(f"Data integrity check FAILED: hash mismatch") + logger.warning("Data integrity check FAILED: hash mismatch") return False @@ -160,44 +161,54 @@ def __init__( self, compliance_level: ComplianceLevel = ComplianceLevel.STRICT, # Stricter thresholds for regulatory compliance (optional overrides) - strict_herg_threshold: Optional[float] = None, - strict_ames_threshold: Optional[float] = None, - strict_hepatotox_threshold: Optional[float] = None, - strict_logp_max: Optional[float] = None, - strict_logp_min: Optional[float] = None, - strict_tpsa_range: Optional[tuple[float, float]] = None, - strict_mw_max: Optional[float] = None, - strict_rotatable_bonds_max: Optional[int] = None, + strict_herg_threshold: float | None = None, + strict_ames_threshold: float | None = None, + strict_hepatotox_threshold: float | None = None, + strict_logp_max: float | None = None, + strict_logp_min: float | None = None, + strict_tpsa_range: tuple[float, float] | None = None, + strict_mw_max: float | None = None, + strict_rotatable_bonds_max: int | None = None, # Risk tier adjustment factors aromatic_ring_penalty: float = 0.15, # Per aromatic ring basic_amine_penalty: float = 0.20, # Per basic nitrogen halogenation_penalty: float = 0.10, # Per halogen atom ): """Initialize strict compliance gate with parametrized thresholds. - + Thresholds are adjusted based on compliance_level if not explicitly provided. """ self.compliance_level = compliance_level - + # Set thresholds based on compliance level if not explicitly provided level_config = self._get_thresholds_for_level(compliance_level) - - self.strict_herg_threshold = strict_herg_threshold if strict_herg_threshold is not None else level_config["herg"] - self.strict_ames_threshold = strict_ames_threshold if strict_ames_threshold is not None else level_config["ames"] - self.strict_hepatotox_threshold = strict_hepatotox_threshold if strict_hepatotox_threshold is not None else level_config["hepatotox"] + + self.strict_herg_threshold = ( + strict_herg_threshold if strict_herg_threshold is not None else level_config["herg"] + ) + self.strict_ames_threshold = ( + strict_ames_threshold if strict_ames_threshold is not None else level_config["ames"] + ) + self.strict_hepatotox_threshold = ( + strict_hepatotox_threshold if strict_hepatotox_threshold is not None else level_config["hepatotox"] + ) self.strict_logp_max = strict_logp_max if strict_logp_max is not None else level_config["logp_max"] self.strict_logp_min = strict_logp_min if strict_logp_min is not None else level_config["logp_min"] - + tpsa_range = strict_tpsa_range if strict_tpsa_range is not None else level_config["tpsa_range"] self.strict_tpsa_min, self.strict_tpsa_max = tpsa_range - + self.strict_mw_max = strict_mw_max if strict_mw_max is not None else level_config["mw_max"] - self.strict_rotatable_bonds_max = strict_rotatable_bonds_max if strict_rotatable_bonds_max is not None else level_config["rotatable_bonds_max"] - + self.strict_rotatable_bonds_max = ( + strict_rotatable_bonds_max + if strict_rotatable_bonds_max is not None + else level_config["rotatable_bonds_max"] + ) + self.aromatic_ring_penalty = aromatic_ring_penalty self.basic_amine_penalty = basic_amine_penalty self.halogenation_penalty = halogenation_penalty - + def _get_thresholds_for_level(self, level: ComplianceLevel) -> dict[str, Any]: """Get default thresholds for compliance level.""" if level == ComplianceLevel.RELAXED: @@ -248,54 +259,47 @@ def _get_thresholds_for_level(self, level: ComplianceLevel) -> dict[str, Any]: def evaluate( self, smiles: str, - toxicity_probs: Optional[dict[str, float]] = None, + toxicity_probs: dict[str, float] | None = None, user_id: str = "system", ) -> QualityAssessment: """Evaluate molecule against strict compliance criteria. - + Args: smiles: SMILES string toxicity_probs: Pre-computed toxicity probabilities (optional) user_id: User performing evaluation (for audit trail) - + Returns: QualityAssessment with comprehensive compliance evaluation """ # Generate audit ID for traceability audit_id = self._generate_audit_id(smiles, user_id) - + # Validate SMILES if not self._validate_smiles(smiles): return self._create_rejection_assessment( - smiles, - audit_id, - "Invalid SMILES structure", - "Failed basic SMILES validation" + smiles, audit_id, "Invalid SMILES structure", "Failed basic SMILES validation" ) - + # Calculate molecular properties props = self._calculate_properties(smiles) - + # Identify risk factors risk_factors = self._identify_risk_factors(props) - + # Run compliance checks checks = self._run_compliance_checks(props, toxicity_probs or {}) - + # Calculate quality tier - quality_tier, tier_confidence = self._classify_quality_tier( - checks, risk_factors, props - ) - + quality_tier, tier_confidence = self._classify_quality_tier(checks, risk_factors, props) + # Verify data integrity integrity = self._verify_data_integrity(smiles, props) - + # Generate assessment overall_passed = all(c.passed for c in checks) - recommendation = self._generate_recommendation( - quality_tier, checks, risk_factors - ) - + recommendation = self._generate_recommendation(quality_tier, checks, risk_factors) + return QualityAssessment( smiles=smiles, quality_tier=quality_tier, @@ -316,13 +320,13 @@ def _validate_smiles(self, smiles: str) -> bool: """Validate SMILES string structure.""" if not isinstance(smiles, str) or not smiles.strip(): return False - + # Check for obviously invalid characters or patterns invalid_chars = set("!@#$%&*()[]{}\\/<>|~`") if any(c in smiles for c in invalid_chars if c not in "[]()"): # Note: [] and () are valid in SMILES return False - + # If RDKit is available, use it for proper validation if _RDKIT: mol = Chem.MolFromSmiles(smiles) @@ -336,18 +340,15 @@ def _validate_smiles(self, smiles: str) -> bool: except Exception: return False return True - + # Fallback: at least check basic SMILES character set # Valid SMILES characters (simplified) valid_chars = set("CNOPSFClBrIF=][()\\@+#-0123456789%") if not all(c in valid_chars for c in smiles if c not in "clno"): return False - + # Must contain at least one element symbol or number - if not any(c in smiles for c in "CNOPSFClBr0123456789"): - return False - - return True + return any(c in smiles for c in "CNOPSFClBr0123456789") def _calculate_properties(self, smiles: str) -> dict[str, float]: """Calculate molecular properties.""" @@ -366,7 +367,7 @@ def _calculate_properties(self, smiles: str) -> dict[str, float]: "halogen_count": self._count_halogens(mol), "heavy_atoms": int(mol.GetNumHeavyAtoms()), } - + # Fallback heuristic return self._estimate_properties_heuristic(smiles) @@ -374,10 +375,7 @@ def _count_aromatic_rings(self, mol: Any) -> int: """Count aromatic rings.""" try: ri = mol.GetRingInfo() - return sum( - 1 for ring in ri.AtomRings() - if all(mol.GetAtomWithIdx(i).GetIsAromatic() for i in ring) - ) + return sum(1 for ring in ri.AtomRings() if all(mol.GetAtomWithIdx(i).GetIsAromatic() for i in ring)) except Exception: return 0 @@ -386,9 +384,8 @@ def _count_basic_nitrogens(self, mol: Any) -> int: count = 0 try: for atom in mol.GetAtoms(): - if atom.GetSymbol() == "N": - if atom.GetTotalDegree() < 3: # Free lone pair - count += 1 + if atom.GetSymbol() == "N" and atom.GetTotalDegree() < 3: # Free lone pair + count += 1 except Exception: pass return count @@ -424,25 +421,25 @@ def _estimate_properties_heuristic(self, smiles: str) -> dict[str, float]: def _identify_risk_factors(self, props: dict[str, float]) -> list[RiskFactor]: """Identify known risk factors.""" risk_factors = [] - + if props.get("logp", 0) > 3.5: risk_factors.append(RiskFactor.HIGH_LOGP) - + if props.get("mw", 0) > 450: risk_factors.append(RiskFactor.HIGH_MW) - + if props.get("aromatic_rings", 0) >= 3: risk_factors.append(RiskFactor.AROMATIC_TOXICOPHORE) - + if props.get("basic_nitrogens", 0) >= 2: risk_factors.append(RiskFactor.BASIC_AMINE) - + if props.get("halogen_count", 0) > 0: risk_factors.append(RiskFactor.HALOGENATION) - + if props.get("tpsa", 0) < 20: risk_factors.append(RiskFactor.POOR_SOLUBILITY) - + return risk_factors def _run_compliance_checks( @@ -452,7 +449,7 @@ def _run_compliance_checks( ) -> list[ComplianceCheckResult]: """Run all compliance checks.""" checks = [] - + # LogP check logp = props.get("logp", 0) checks.append( @@ -464,10 +461,12 @@ def _run_compliance_checks( message=f"LogP = {logp:.2f}, allowed range [{self.strict_logp_min}, {self.strict_logp_max}]", value=logp, threshold=self.strict_logp_max, - remediation="Reduce lipophilicity by removing hydrophobic groups" if logp > self.strict_logp_max else None, + remediation=( + "Reduce lipophilicity by removing hydrophobic groups" if logp > self.strict_logp_max else None + ), ) ) - + # Molecular weight check mw = props.get("mw", 0) checks.append( @@ -481,7 +480,7 @@ def _run_compliance_checks( threshold=self.strict_mw_max, ) ) - + # TPSA check tpsa = props.get("tpsa", 0) checks.append( @@ -495,7 +494,7 @@ def _run_compliance_checks( threshold=self.strict_tpsa_max, ) ) - + # Rotatable bonds check rot_bonds = props.get("rot_bonds", 0) checks.append( @@ -509,7 +508,7 @@ def _run_compliance_checks( threshold=float(self.strict_rotatable_bonds_max), ) ) - + # Toxicity endpoint checks if "herg" in toxicity_probs: herg_prob = toxicity_probs["herg"] @@ -524,7 +523,7 @@ def _run_compliance_checks( threshold=self.strict_herg_threshold, ) ) - + if "ames" in toxicity_probs: ames_prob = toxicity_probs["ames"] checks.append( @@ -538,7 +537,7 @@ def _run_compliance_checks( threshold=self.strict_ames_threshold, ) ) - + if "hepatotox" in toxicity_probs: hep_prob = toxicity_probs["hepatotox"] checks.append( @@ -552,7 +551,7 @@ def _run_compliance_checks( threshold=self.strict_hepatotox_threshold, ) ) - + return checks def _classify_quality_tier( @@ -564,15 +563,15 @@ def _classify_quality_tier( """Classify quality tier based on compliance and risk.""" failed_critical = sum(1 for c in checks if not c.passed and c.severity == "critical") failed_major = sum(1 for c in checks if not c.passed and c.severity == "major") - + # Risk factor penalty risk_penalty = 0.0 risk_penalty += len([r for r in risk_factors if r == RiskFactor.HIGH_LOGP]) * self.aromatic_ring_penalty risk_penalty += len([r for r in risk_factors if r == RiskFactor.BASIC_AMINE]) * self.basic_amine_penalty risk_penalty += len([r for r in risk_factors if r == RiskFactor.HALOGENATION]) * self.halogenation_penalty - + confidence = max(0.5, 1.0 - risk_penalty - failed_critical * 0.3 - failed_major * 0.1) - + if failed_critical > 0: return QualityTier.REJECTED, confidence elif failed_major > 0: @@ -598,16 +597,16 @@ def _verify_data_integrity( } data_json = json.dumps(data, sort_keys=True, separators=(",", ":")) checksum = hashlib.sha256(data_json.encode()).hexdigest() - + report = DataIntegrityReport( checksum=checksum, timestamp=datetime.utcnow(), data_hash=hashlib.sha256(smiles.encode()).hexdigest(), ) - + # Integrity is automatically verified by checksum calculation report.integrity_verified = True - + return report def _generate_recommendation( @@ -621,17 +620,17 @@ def _generate_recommendation( failed = [c for c in checks if not c.passed] reasons = ", ".join(c.check_name for c in failed[:2]) return f"REJECTED - Failed critical checks: {reasons}" - + elif quality_tier == QualityTier.TIER_1: return "APPROVED - Tier 1 (excellent safety profile, ready for IND submission)" - + elif quality_tier == QualityTier.TIER_2: rf_str = ", ".join(r.value for r in risk_factors[:2]) return f"APPROVED - Tier 2 (acceptable with standard monitoring; risk factors: {rf_str})" - + elif quality_tier == QualityTier.TIER_3: return "CONDITIONAL - Tier 3 (requires enhanced preclinical evaluation)" - + else: # TIER_4 return "MARGINAL - Tier 4 (additional studies required; not recommended for IND at this time)" @@ -655,13 +654,13 @@ def _create_rejection_assessment( severity="critical", message=message, ) - + integrity = DataIntegrityReport( checksum="", timestamp=datetime.utcnow(), data_hash="", ) - + return QualityAssessment( smiles=smiles, quality_tier=QualityTier.REJECTED, @@ -680,36 +679,38 @@ def evaluate_batch_with_strict_compliance( compliance_level: ComplianceLevel = ComplianceLevel.STRICT, ) -> dict[str, Any]: """Batch evaluation with strict compliance. - + Args: smiles_list: List of SMILES strings compliance_level: Regulatory compliance level - + Returns: Dictionary with batch results and compliance summary """ gate = StrictComplianceGate(compliance_level=compliance_level) - + assessments = {} passed_count = 0 rejected_count = 0 critical_issues = [] - + for smiles in smiles_list: assessment = gate.evaluate(smiles) assessments[smiles] = assessment - + if assessment.overall_passed: passed_count += 1 else: rejected_count += 1 - critical_issues.append({ - "smiles": smiles, - "audit_id": assessment.audit_id, - "tier": assessment.quality_tier.value, - "recommendation": assessment.recommendation, - }) - + critical_issues.append( + { + "smiles": smiles, + "audit_id": assessment.audit_id, + "tier": assessment.quality_tier.value, + "recommendation": assessment.recommendation, + } + ) + return { "compliance_level": compliance_level.value, "total_evaluated": len(smiles_list), @@ -717,8 +718,5 @@ def evaluate_batch_with_strict_compliance( "rejected": rejected_count, "pass_rate": passed_count / max(1, len(smiles_list)), "critical_issues": critical_issues, - "assessments": { - smiles: assessment.as_dict() - for smiles, assessment in assessments.items() - }, + "assessments": {smiles: assessment.as_dict() for smiles, assessment in assessments.items()}, } diff --git a/drug_discovery/safety/toxicity_gate.py b/drug_discovery/safety/toxicity_gate.py index 54726e3f71..f584db0bbf 100644 --- a/drug_discovery/safety/toxicity_gate.py +++ b/drug_discovery/safety/toxicity_gate.py @@ -81,10 +81,7 @@ def as_dict(self) -> dict[str, Any]: "safety_score": self.safety_score, "drug_likeness": self.drug_likeness, "lipinski_violations": self.lipinski_violations, - "endpoints": [ - {"name": e.name, "probability": e.probability, "passed": e.passed} - for e in self.endpoints - ], + "endpoints": [{"name": e.name, "probability": e.probability, "passed": e.passed} for e in self.endpoints], "rejection_reasons": self.rejection_reasons, } @@ -176,7 +173,9 @@ def _evaluate_internal(self, smiles: str) -> ToxicityVerdict: # Hepatotoxicity hepato_p = admet_scores.get("hepatotox", self._estimate_hepatotox(props)) ep = EndpointScore( - "Hepatotoxicity", hepato_p, self.config.hepatotox_threshold, + "Hepatotoxicity", + hepato_p, + self.config.hepatotox_threshold, hepato_p <= self.config.hepatotox_threshold, ) endpoints.append(ep) @@ -186,7 +185,9 @@ def _evaluate_internal(self, smiles: str) -> ToxicityVerdict: # Cytotoxicity cyto_p = admet_scores.get("cytotox", self._estimate_cytotox(props)) ep = EndpointScore( - "Cytotoxicity", cyto_p, self.config.cytotox_threshold, + "Cytotoxicity", + cyto_p, + self.config.cytotox_threshold, cyto_p <= self.config.cytotox_threshold, ) endpoints.append(ep) @@ -221,7 +222,7 @@ def _evaluate_internal(self, smiles: str) -> ToxicityVerdict: drug_likeness=drug_likeness, lipinski_violations=lipinski_violations, rejection_reasons=rejection_reasons, - metadata={"suggested_counters": suggested_counters} + metadata={"suggested_counters": suggested_counters}, ) def _suggest_counter_toxins(self, endpoints: list[EndpointScore]) -> list[str]: @@ -370,7 +371,7 @@ def _compute_drug_likeness(props: dict[str, float], toxicity: float) -> float: # Geometric mean product = d_mw * d_logp * d_hba * d_hbd * d_tox - return max(0.0, product ** 0.2) + return max(0.0, product**0.2) # --------------------------------------------------------------------------- diff --git a/drug_discovery/screening/__init__.py b/drug_discovery/screening/__init__.py index 1ec67f7784..183099d9a7 100644 --- a/drug_discovery/screening/__init__.py +++ b/drug_discovery/screening/__init__.py @@ -1,2 +1,2 @@ from .admet_models import ADMETScreen -from .filtering import filter_admet \ No newline at end of file +from .filtering import filter_admet diff --git a/drug_discovery/screening/admet_models.py b/drug_discovery/screening/admet_models.py index 6b892d3b66..301ef4425e 100644 --- a/drug_discovery/screening/admet_models.py +++ b/drug_discovery/screening/admet_models.py @@ -1,48 +1,49 @@ -from typing import List, Dict, Any -from ..data.rdkit_utils import smiles_to_mols, compute_descriptors - from drug_discovery.glp_tox_panel import PreClinicalToxPanel +from ..data.rdkit_utils import compute_descriptors, smiles_to_mols + + class ADMETScreen: def __init__(self): try: import deepchem as dc + self.deepchem_available = True # Enhanced Tox21 + Clintox for hepatotox - self.tox21_tasks = ['SR-HSE', 'SR-p53'] # proxies - self.tox21_model = dc.models.GraphConvModel(n_tasks=2, mode='classification') # Stub/train on Tox21 - self.clintox_model = dc.models.GraphConvModel(n_tasks=2, mode='classification') # Hepatotox from Clintox + self.tox21_tasks = ["SR-HSE", "SR-p53"] # proxies + self.tox21_model = dc.models.GraphConvModel(n_tasks=2, mode="classification") # Stub/train on Tox21 + self.clintox_model = dc.models.GraphConvModel(n_tasks=2, mode="classification") # Hepatotox from Clintox self.featurizer = dc.feat.CircularFingerprint(size=1024) except: self.deepchem_available = False - + self.glp_panel = PreClinicalToxPanel(herg_threshold=0.3) # Strict CiPA-like hERG - def predict(self, smiles_list: List[str]) -> Dict[str, List[float]]: + def predict(self, smiles_list: list[str]) -> dict[str, list[float]]: mols = smiles_to_mols(smiles_list) - results = {'herg': [], 'hepatotox': [], 'bbb': [], 'qed': []} - + results = {"herg": [], "hepatotox": [], "bbb": [], "qed": []} + # Enhanced hERG from GLP panel (heuristic + pharmacophore) for smi in smiles_list: panel = self.glp_panel.evaluate(smi) - results['herg'].append(panel.herg.inhibition_probability) - + results["herg"].append(panel.herg.inhibition_probability) + # DeepChem stubs for others if self.deepchem_available: X = self.featurizer.featurize(smiles_list) - tox21_pred = self.tox21_model.predict(X) + self.tox21_model.predict(X) clintox_pred = self.clintox_model.predict(X) for i in range(len(smiles_list)): - results['hepatotox'].append(clintox_pred[i]['probabilities'][1]) # hepatotox class - + results["hepatotox"].append(clintox_pred[i]["probabilities"][1]) # hepatotox class + else: - results['hepatotox'] = [0.05] * len(smiles_list) - results['bbb'] = [0.8] * len(smiles_list) - + results["hepatotox"] = [0.05] * len(smiles_list) + results["bbb"] = [0.8] * len(smiles_list) + desc_df = compute_descriptors(mols) - results['qed'] = desc_df['qed'].tolist() - + results["qed"] = desc_df["qed"].tolist() + # CiPA-like hERG risk bands - results['herg_risk'] = ['high' if p > 0.5 else 'moderate' if p > 0.3 else 'low' for p in results['herg']] - + results["herg_risk"] = ["high" if p > 0.5 else "moderate" if p > 0.3 else "low" for p in results["herg"]] + return results diff --git a/drug_discovery/screening/filtering.py b/drug_discovery/screening/filtering.py index 67506d5dcc..ffdaedcc0c 100644 --- a/drug_discovery/screening/filtering.py +++ b/drug_discovery/screening/filtering.py @@ -1,11 +1,11 @@ -def filter_admet(predictions: dict, thresholds: dict = None) -> list: +def filter_admet(predictions: dict, thresholds: dict | None = None) -> list: if thresholds is None: - thresholds = {'herg': 0.3, 'hepatotox': 0.1, 'qed': 0.5} # Strict hERG CiPA (p<0.3 low risk) + thresholds = {"herg": 0.3, "hepatotox": 0.1, "qed": 0.5} # Strict hERG CiPA (p<0.3 low risk) filtered_indices = [] - for i in range(len(predictions['herg'])): - herg_ok = predictions['herg'][i] < thresholds['herg'] - hepat_ok = predictions['hepatotox'][i] < thresholds['hepatotox'] - qed_ok = predictions['qed'][i] > thresholds['qed'] + for i in range(len(predictions["herg"])): + herg_ok = predictions["herg"][i] < thresholds["herg"] + hepat_ok = predictions["hepatotox"][i] < thresholds["hepatotox"] + qed_ok = predictions["qed"][i] > thresholds["qed"] if herg_ok and hepat_ok and qed_ok: filtered_indices.append(i) - return filtered_indices \ No newline at end of file + return filtered_indices diff --git a/drug_discovery/simulation/__init__.py b/drug_discovery/simulation/__init__.py index e0aa95d176..87ce8d651a 100644 --- a/drug_discovery/simulation/__init__.py +++ b/drug_discovery/simulation/__init__.py @@ -8,12 +8,12 @@ from .patient_generator import PatientGenerator as PatientGenerator __all__ = [ - "CGSimulator", - "PatientGenerator", "BayesianPKPD", + "CGSimulator", "ClinicalTrialSimulator", "MicrogravitySimulator", "OrbitalLogisticsOptimizer", + "PatientGenerator", ] try: @@ -30,6 +30,6 @@ generate_lambda_schedule as generate_lambda_schedule, ) - __all__.extend(["FEPPipeline", "FEPConfig", "FEPSurrogateNetwork", "generate_lambda_schedule"]) + __all__.extend(["FEPConfig", "FEPPipeline", "FEPSurrogateNetwork", "generate_lambda_schedule"]) except ImportError: pass diff --git a/drug_discovery/simulation/clinical_trial.py b/drug_discovery/simulation/clinical_trial.py index e4c8a25b1a..0ea340f6cb 100644 --- a/drug_discovery/simulation/clinical_trial.py +++ b/drug_discovery/simulation/clinical_trial.py @@ -96,14 +96,14 @@ def simulate_phase3( # Calculate p-value using two-sample z-test n_treatment = len(treatment_results) n_control = len(control_results) - + # Pooled standard error for proportion difference p_pooled = (treatment_results.sum() + control_results.sum()) / (n_treatment + n_control) - se = np.sqrt(p_pooled * (1 - p_pooled) * (1/n_treatment + 1/n_control)) - + se = np.sqrt(p_pooled * (1 - p_pooled) * (1 / n_treatment + 1 / n_control)) + # Z-statistic z_stat = (treatment_rate - control_rate) / (se + 1e-10) # Add small epsilon to avoid division by zero - + # Two-tailed p-value from standard normal distribution # P(|Z| > |z_stat|) = 2 * P(Z > |z_stat|) p_value = 2 * (1 - self._standard_normal_cdf(abs(z_stat))) diff --git a/drug_discovery/simulation/coarse_grained_md.py b/drug_discovery/simulation/coarse_grained_md.py index aac2f5c312..57ba1e685b 100644 --- a/drug_discovery/simulation/coarse_grained_md.py +++ b/drug_discovery/simulation/coarse_grained_md.py @@ -36,24 +36,24 @@ def run_simulation(self, system_name: str, steps: int = 10000) -> dict[str, any] In reality, this would interface with GROMACS or OpenMM using Martini forcefield. """ logger.info(f"Running CG simulation for {system_name} for {steps} steps") - + # Heuristic potential energy based on system size and steps # Larger systems and more steps -> more relaxation -> lower potential energy base_energy = -1000.0 system_factor = len(system_name) * 10 # System complexity factor convergence_factor = -1.0 * np.sqrt(steps / 1000) # Energy relaxation over time potential_energy = base_energy - system_factor + convergence_factor * np.random.normal(0, 50) - + # Diffusion coefficient decreases with larger/more ordered systems # Ranges from 1e-8 to 1e-6 m^2/s (typical biological scales) base_diffusion = 0.1 * np.exp(-len(system_name) / 20) diffusion_coefficient = base_diffusion * (1.0 + 0.1 * np.random.normal(0, 1)) - + # Convergence check: did we reach reasonable stability? # Systems with fewer atoms and more steps converge better convergence_score = min(0.95, (steps / 10000) * (1.0 / (1.0 + len(system_name) / 100))) converged = convergence_score > 0.7 - + return { "status": "completed", "final_potential_energy": float(potential_energy), @@ -72,7 +72,7 @@ def analyze_aggregation(self) -> float: # When full MDAnalysis is available: # contacts_analysis = contacts.Contacts(...) # aggregation_index = contacts_analysis.count() - + try: # If we have trajectory data, estimate aggregation from Rg changes rg_values = self.calculate_radius_of_gyration() @@ -87,7 +87,7 @@ def analyze_aggregation(self) -> float: return float(aggregation_index) except Exception: pass - + # Fallback: estimate from number of atoms (more atoms -> higher aggregation potential) try: if self.universe: @@ -97,7 +97,7 @@ def analyze_aggregation(self) -> float: return float(aggregation_index) except Exception: pass - + # Default heuristic return 0.65 diff --git a/drug_discovery/simulation/free_energy.py b/drug_discovery/simulation/free_energy.py index 0bc6d0b194..0a201a8768 100644 --- a/drug_discovery/simulation/free_energy.py +++ b/drug_discovery/simulation/free_energy.py @@ -142,10 +142,16 @@ def predict_dd_g(self, ligand_a: dict[str, Any], ligand_b: dict[str, Any]) -> di for lv in self.lambdas: lam = torch.tensor([lv], device=self.device, dtype=torch.float32) # type: ignore[union-attr] ea = self.surrogate( - ligand_a["z"].to(self.device), ligand_a["pos"].to(self.device), ligand_a["edges"].to(self.device), lam + ligand_a["z"].to(self.device), + ligand_a["pos"].to(self.device), + ligand_a["edges"].to(self.device), + lam, ).item() eb = self.surrogate( - ligand_b["z"].to(self.device), ligand_b["pos"].to(self.device), ligand_b["edges"].to(self.device), lam + ligand_b["z"].to(self.device), + ligand_b["pos"].to(self.device), + ligand_b["edges"].to(self.device), + lam, ).item() energies_a.append(ea) energies_b.append(eb) diff --git a/drug_discovery/smd/abfe_residuals.py b/drug_discovery/smd/abfe_residuals.py index 9f9f09f094..a5cee72d01 100644 --- a/drug_discovery/smd/abfe_residuals.py +++ b/drug_discovery/smd/abfe_residuals.py @@ -3,16 +3,18 @@ Provides residual computations and outlier identification to track simulation convergence and identify problematic systems. """ + from __future__ import annotations -from typing import Sequence, List, Dict, Any import math +from collections.abc import Sequence +from typing import Any -def compute_residuals(predicted: Sequence[float], observed: Sequence[float]) -> List[float]: +def compute_residuals(predicted: Sequence[float], observed: Sequence[float]) -> list[float]: if len(predicted) != len(observed): raise ValueError("predicted and observed length mismatch") - return [p - o for p, o in zip(predicted, observed)] + return [p - o for p, o in zip(predicted, observed, strict=False)] def rmse(residuals: Sequence[float]) -> float: @@ -21,7 +23,7 @@ def rmse(residuals: Sequence[float]) -> float: return math.sqrt(sum(r * r for r in residuals) / len(residuals)) -def z_scores(residuals: Sequence[float]) -> List[float]: +def z_scores(residuals: Sequence[float]) -> list[float]: n = len(residuals) if n == 0: return [] @@ -33,12 +35,12 @@ def z_scores(residuals: Sequence[float]) -> List[float]: return [(r - mean) / sd for r in residuals] -def identify_outliers(residuals: Sequence[float], threshold_z: float = 2.5) -> List[int]: +def identify_outliers(residuals: Sequence[float], threshold_z: float = 2.5) -> list[int]: zs = z_scores(residuals) return [i for i, z in enumerate(zs) if abs(z) >= threshold_z] -def summarize_abfe(predicted: Sequence[float], observed: Sequence[float], top_n: int = 5) -> Dict[str, Any]: +def summarize_abfe(predicted: Sequence[float], observed: Sequence[float], top_n: int = 5) -> dict[str, Any]: res = compute_residuals(predicted, observed) summary = { "n": len(res), diff --git a/drug_discovery/strategy/__init__.py b/drug_discovery/strategy/__init__.py index 8ff8369a6e..058a682f47 100644 --- a/drug_discovery/strategy/__init__.py +++ b/drug_discovery/strategy/__init__.py @@ -6,9 +6,9 @@ __all__ = [ "CandidateProfile", - "TargetProductProfile", - "TPPScorer", "ManufacturingPlan", "ManufacturingStrategyPlanner", "ProgramStrategyEngine", + "TPPScorer", + "TargetProductProfile", ] diff --git a/drug_discovery/structure_analysis/__init__.py b/drug_discovery/structure_analysis/__init__.py index 02ec8a1379..5b5f6753c8 100644 --- a/drug_discovery/structure_analysis/__init__.py +++ b/drug_discovery/structure_analysis/__init__.py @@ -1,2 +1,2 @@ from .cif_parser import parse_cif_to_mol -from .xrpd_analysis import analyze_xrpd \ No newline at end of file +from .xrpd_analysis import analyze_xrpd diff --git a/drug_discovery/structure_analysis/cif_parser.py b/drug_discovery/structure_analysis/cif_parser.py index 0b6e168e4f..df2005a0f9 100644 --- a/drug_discovery/structure_analysis/cif_parser.py +++ b/drug_discovery/structure_analysis/cif_parser.py @@ -1,11 +1,13 @@ try: import gemmi + GEMMI_AVAILABLE = True except ImportError: GEMMI_AVAILABLE = False from rdkit import Chem + def parse_cif_to_mol(cif_path: str) -> Chem.Mol: if not GEMMI_AVAILABLE: raise ImportError("gemmi required for CIF parsing") @@ -13,4 +15,4 @@ def parse_cif_to_mol(cif_path: str) -> Chem.Mol: # Convert first model/chain to PDB string pdb_block = structure.make_pdb_block().as_string() mol = Chem.MolFromPDBBlock(pdb_block) - return mol \ No newline at end of file + return mol diff --git a/drug_discovery/structure_analysis/xrpd_analysis.py b/drug_discovery/structure_analysis/xrpd_analysis.py index f9d4c006ad..405467385a 100644 --- a/drug_discovery/structure_analysis/xrpd_analysis.py +++ b/drug_discovery/structure_analysis/xrpd_analysis.py @@ -1,14 +1,16 @@ -from scipy.signal import find_peaks import numpy as np from rdkit.Chem import Descriptors +from scipy.signal import find_peaks + def simulate_xrpd_pattern(mol): # Stub: simple peak simulation based on unit cell guess from descriptors two_theta = np.linspace(5, 50, 1000) intensity = np.random.rand(1000) * Descriptors.MolWt(mol) / 500 # Mock peaks, _ = find_peaks(intensity, height=0.1 * np.max(intensity)) - return {'two_theta': two_theta[peaks].tolist(), 'intensity': intensity[peaks].tolist()} + return {"two_theta": two_theta[peaks].tolist(), "intensity": intensity[peaks].tolist()} + def analyze_xrpd(xrpd_file: str, mol): pattern = simulate_xrpd_pattern(mol) - return pattern \ No newline at end of file + return pattern diff --git a/drug_discovery/synthesis/__init__.py b/drug_discovery/synthesis/__init__.py index 23c88129a8..4c57f06258 100644 --- a/drug_discovery/synthesis/__init__.py +++ b/drug_discovery/synthesis/__init__.py @@ -8,15 +8,15 @@ from .retrosynthesis import RetrosynthesisPlanner, SynthesisFeasibilityScorer __all__ = [ - "RetrosynthesisPlanner", - "SynthesisFeasibilityScorer", "AiZynthFinderBackend", "BackendResult", "BaseRetrosynthesisBackend", - "RouteCandidate", "MolecularTransformerAdapter", - "ReactionPrediction", - "PistachioDatasets", "PistachioDatasetResult", + "PistachioDatasets", + "ReactionPrediction", "ReactionRecord", + "RetrosynthesisPlanner", + "RouteCandidate", + "SynthesisFeasibilityScorer", ] diff --git a/drug_discovery/synthesis/retrosynthesis.py b/drug_discovery/synthesis/retrosynthesis.py index 0efd418d87..2a31c26aeb 100644 --- a/drug_discovery/synthesis/retrosynthesis.py +++ b/drug_discovery/synthesis/retrosynthesis.py @@ -4,9 +4,9 @@ # pyright: reportMissingTypeStubs=false, reportUnknownMemberType=false, reportUnknownArgumentType=false +import concurrent.futures import logging import os -import concurrent.futures from collections.abc import Sequence import numpy as np @@ -43,6 +43,7 @@ def __init__( self.use_ray = use_ray if self.use_ray: import ray + if not ray.is_initialized(): ray.init(ignore_reinit_error=True, address=os.getenv("RAY_ADDRESS")) @@ -62,10 +63,9 @@ def _run_backends(self, target_smiles: str, max_depth: int) -> tuple[RouteCandid with concurrent.futures.ThreadPoolExecutor(max_workers=len(self.backends)) as executor: future_to_backend = { - executor.submit(backend.plan, target_smiles, max_depth=max_depth): backend - for backend in self.backends + executor.submit(backend.plan, target_smiles, max_depth=max_depth): backend for backend in self.backends } - + for future in concurrent.futures.as_completed(future_to_backend): backend = future_to_backend[future] try: @@ -84,7 +84,7 @@ def _run_backends(self, target_smiles: str, max_depth: int) -> tuple[RouteCandid r.score if r.score is not None else float("inf"), ), )[0] - + if route_choice is None: route_choice = best else: @@ -93,7 +93,7 @@ def _run_backends(self, target_smiles: str, max_depth: int) -> tuple[RouteCandid curr_score = route_choice.score if route_choice.score is not None else float("inf") new_steps = best.steps if best.steps is not None else max_depth + 10 new_score = best.score if best.score is not None else float("inf") - + if (new_steps, new_score) < (curr_steps, curr_score): route_choice = best @@ -152,10 +152,11 @@ def plan_synthesis_batch(self, smiles_list: list[str], max_depth: int = 5) -> li """Plan synthesis for a batch of molecules, optionally using Ray.""" if self.use_ray: import ray + @ray.remote def remote_plan(planner, smiles, depth): return planner.plan_synthesis(smiles, max_depth=depth) - + futures = [remote_plan.remote(self, smiles, max_depth) for smiles in smiles_list] return ray.get(futures) else: @@ -245,9 +246,9 @@ def score_synthetic_accessibility(self, smiles: str) -> float: # Use RDKit's built-in SA score if available # Otherwise, use simple heuristics - ring_count = getattr(Descriptors, "RingCount") - num_rotatable_bonds = getattr(Descriptors, "NumRotatableBonds") - mol_wt_func = getattr(Descriptors, "MolWt") + ring_count = Descriptors.RingCount + num_rotatable_bonds = Descriptors.NumRotatableBonds + mol_wt_func = Descriptors.MolWt num_rings = ring_count(mol) num_stereo = len(Chem.FindMolChiralCenters(mol, includeUnassigned=True)) @@ -303,16 +304,16 @@ def score_feasibility(self, smiles: str, retro_plan: dict | None = None) -> dict scores["sa_score"] = 11 - sa_score # Convert to 1-10, higher better # Molecular complexity - ring_count = getattr(Descriptors, "RingCount") - num_heteroatoms_func = getattr(Lipinski, "NumHeteroatoms") + ring_count = Descriptors.RingCount + num_heteroatoms_func = Lipinski.NumHeteroatoms num_rings = ring_count(mol) num_heteroatoms = num_heteroatoms_func(mol) complexity = (num_rings + num_heteroatoms) / 10.0 scores["complexity"] = max(0, 10 - complexity * 10) # Functional group diversity (penalize exotic groups) - num_aliphatic_rings = getattr(Descriptors, "NumAliphaticRings") - num_aromatic_rings = getattr(Descriptors, "NumAromaticRings") + num_aliphatic_rings = Descriptors.NumAliphaticRings + num_aromatic_rings = Descriptors.NumAromaticRings num_functional_groups = num_aliphatic_rings(mol) + num_aromatic_rings(mol) scores["fg_score"] = min(10, num_functional_groups) diff --git a/drug_discovery/testing/__init__.py b/drug_discovery/testing/__init__.py index e9f5ce6164..3570112554 100644 --- a/drug_discovery/testing/__init__.py +++ b/drug_discovery/testing/__init__.py @@ -16,8 +16,8 @@ from drug_discovery.testing.uncertainty import UncertaintyEstimator __all__ = [ - "ToxicityPredictor", "DrugCombinationTester", "RobustnessTester", + "ToxicityPredictor", "UncertaintyEstimator", ] diff --git a/drug_discovery/testing/cns_bbb_penetration.py b/drug_discovery/testing/cns_bbb_penetration.py index 1f608918d0..07940e5bd5 100644 --- a/drug_discovery/testing/cns_bbb_penetration.py +++ b/drug_discovery/testing/cns_bbb_penetration.py @@ -1,7 +1,9 @@ +from typing import Any + import numpy as np from rdkit import Chem from rdkit.Chem import Descriptors, rdMolDescriptors -from typing import Dict, Any + class BioavailabilityScreener: """ @@ -26,24 +28,35 @@ def calculate_cns_mpo(self, smiles: str) -> float: mw = Descriptors.MolWt(mol) tpsa = rdMolDescriptors.CalcTPSA(mol) hbd = rdMolDescriptors.CalcNumHBD(mol) - + # pKa estimation (very heuristic for this scaffolding) # Bases often have pKa ~9, Acids ~4. - pka = 8.0 # Default fallback - if mol.HasSubstructMatch(Chem.MolFromSmarts("[NX3;H2,H1,H0;c,C]")): # Amine + pka = 8.0 # Default fallback + if mol.HasSubstructMatch(Chem.MolFromSmarts("[NX3;H2,H1,H0;c,C]")): # Amine pka = 9.5 - elif mol.HasSubstructMatch(Chem.MolFromSmarts("C(=O)[OH]")): # Carboxylic acid + elif mol.HasSubstructMatch(Chem.MolFromSmarts("C(=O)[OH]")): # Carboxylic acid pka = 4.5 # Desirability functions (normalized to [0, 1]) - def f_clogp(x): return 1.0 if x <= 3.0 else np.exp(-(x-3.0)**2 / 2.0) - def f_mw(x): return 1.0 if x <= 360.0 else np.exp(-(x-360.0)**2 / 5000.0) - def f_tpsa(x): - if 40 <= x <= 90: return 1.0 - elif x < 40: return x/40.0 - else: return np.exp(-(x-90.0)**2 / 1000.0) - def f_hbd(x): return 1.0 if x <= 0 else 0.8 if x == 1 else 0.0 - def f_pka(x): return np.exp(-(x-8.0)**2 / 10.0) + def f_clogp(x): + return 1.0 if x <= 3.0 else np.exp(-((x - 3.0) ** 2) / 2.0) + + def f_mw(x): + return 1.0 if x <= 360.0 else np.exp(-((x - 360.0) ** 2) / 5000.0) + + def f_tpsa(x): + if 40 <= x <= 90: + return 1.0 + elif x < 40: + return x / 40.0 + else: + return np.exp(-((x - 90.0) ** 2) / 1000.0) + + def f_hbd(x): + return 1.0 if x <= 0 else 0.8 if x == 1 else 0.0 + + def f_pka(x): + return np.exp(-((x - 8.0) ** 2) / 10.0) scores = [f_clogp(clogp), f_mw(mw), f_tpsa(tpsa), f_hbd(hbd), f_pka(pka)] return float(np.sum(scores)) @@ -51,7 +64,7 @@ def f_pka(x): return np.exp(-(x-8.0)**2 / 10.0) def predict_p_glycoprotein_efflux(self, smiles: str) -> bool: """ Predicts if the drug is a P-gp substrate (active efflux risk). - P-gp substrates are typically large (MW > 400), lipophilic, + P-gp substrates are typically large (MW > 400), lipophilic, and have many H-bond acceptors. """ mol = Chem.MolFromSmiles(smiles) @@ -63,19 +76,16 @@ def predict_p_glycoprotein_efflux(self, smiles: str) -> bool: hba = Descriptors.NumHAcceptors(mol) # Heuristic Rule: P-gp Substrate if MW > 400 and (LogP > 3 or HBA > 8) - if mw > 400 and (logp > 3.0 or hba > 8): - return True - - return False + return bool(mw > 400 and (logp > 3.0 or hba > 8)) - def get_cns_profile(self, smiles: str) -> Dict[str, Any]: + def get_cns_profile(self, smiles: str) -> dict[str, Any]: """Returns the CNS penetration and efflux profile.""" mpo_score = self.calculate_cns_mpo(smiles) pgp_efflux = self.predict_p_glycoprotein_efflux(smiles) - + return { "smiles": smiles, "cns_mpo_score": mpo_score, "pgp_efflux_risk": pgp_efflux, - "likely_cns_penetrant": mpo_score >= 4.0 and not pgp_efflux + "likely_cns_penetrant": mpo_score >= 4.0 and not pgp_efflux, } diff --git a/drug_discovery/testing/drug_combinations.py b/drug_discovery/testing/drug_combinations.py index b020d8e505..5fbcc156f7 100644 --- a/drug_discovery/testing/drug_combinations.py +++ b/drug_discovery/testing/drug_combinations.py @@ -381,7 +381,7 @@ def find_toxicity_balancers( ) -> pd.DataFrame: """ Find molecules that balance the toxicity of a target compound. - + A balancer is a molecule that has an antagonistic effect on the toxicity of another molecule, effectively 'balancing' it. """ @@ -390,16 +390,18 @@ def find_toxicity_balancers( # We use ML synergy model but look for antagonistic interactions on toxicity # In this context, antagonism is good because it reduces the toxic effect. prediction = self.predict_synergy_ml(toxic_smiles, balancer_smiles) - + # If the interaction is antagonistic, it might be a good balancer if prediction["interaction_type"] == "antagonistic": - results.append({ - "toxic_smiles": toxic_smiles, - "balancer_smiles": balancer_smiles, - "balancing_efficiency": abs(prediction["synergy_score"]), - "confidence": prediction["confidence"] - }) - + results.append( + { + "toxic_smiles": toxic_smiles, + "balancer_smiles": balancer_smiles, + "balancing_efficiency": abs(prediction["synergy_score"]), + "confidence": prediction["confidence"], + } + ) + df = pd.DataFrame(results) if not df.empty: df = df.sort_values("balancing_efficiency", ascending=False) diff --git a/drug_discovery/testing/environmental_degradation.py b/drug_discovery/testing/environmental_degradation.py index 3f2336b3a5..1f109d43cc 100644 --- a/drug_discovery/testing/environmental_degradation.py +++ b/drug_discovery/testing/environmental_degradation.py @@ -1,7 +1,9 @@ +from typing import Any + import numpy as np import scipy.constants as const from rdkit import Chem -from typing import Dict, Any + class EnvironmentalStressSimulator: """ @@ -18,31 +20,26 @@ def __init__(self, pre_exponential_factor: float = 1e13): self.A = pre_exponential_factor self.R = const.R # Ideal gas constant in J/(mol*K) - def calculate_arrhenius_decay( - self, - activation_energy_kj: float, - temp_celsius: float, - days: int - ) -> float: + def calculate_arrhenius_decay(self, activation_energy_kj: float, temp_celsius: float, days: int) -> float: """ Calculates the percentage of API remaining using the Arrhenius equation. - + k = A * exp(-Ea / RT) """ # Convert T to Kelvin T = temp_celsius + 273.15 # Convert Ea to Joules Ea = activation_energy_kj * 1000.0 - + # Calculate rate constant k (1/sec) k = self.A * np.exp(-Ea / (self.R * T)) - + # Total seconds of exposure seconds = days * 24 * 3600 - + # Concentration remaining (First order: C = C0 * exp(-kt)) percent_remaining = np.exp(-k * seconds) * 100.0 - + return float(np.clip(percent_remaining, 0.0, 100.0)) def check_photostability(self, smiles: str) -> bool: @@ -61,24 +58,24 @@ def check_photostability(self, smiles: str) -> bool: "nitro_aromatics": "c1ccccc1[N+](=O)[O-]", "phenothiazines": "c1ccc2c(c1)Sc3ccccc3N2", "quinolones": "c1ccc2c(c1)nc(cc2=O)C(=O)O", - "extended_aromatic_systems": "c1ccc2c(c1)ccc3ccccc32" # Pyrene-like + "extended_aromatic_systems": "c1ccc2c(c1)ccc3ccccc32", # Pyrene-like } - for name, smarts in phototoxicity_alerts.items(): + for smarts in phototoxicity_alerts.values(): pattern = Chem.MolFromSmarts(smarts) if pattern and mol.HasSubstructMatch(pattern): return True - + return False - def simulate_stress_report(self, smiles: str, Ea: float, temp: float, days: int) -> Dict[str, Any]: + def simulate_stress_report(self, smiles: str, Ea: float, temp: float, days: int) -> dict[str, Any]: """Comprehensive environmental validation report.""" remaining = self.calculate_arrhenius_decay(Ea, temp, days) photo_risk = self.check_photostability(smiles) - + return { "smiles": smiles, "api_remaining_percent": remaining, "photostability_risk": photo_risk, - "stable_under_conditions": remaining > 90.0 and not photo_risk + "stable_under_conditions": remaining > 90.0 and not photo_risk, } diff --git a/drug_discovery/testing/human_pkpd_dynamics.py b/drug_discovery/testing/human_pkpd_dynamics.py index 5b6971c4ca..53ac75fd01 100644 --- a/drug_discovery/testing/human_pkpd_dynamics.py +++ b/drug_discovery/testing/human_pkpd_dynamics.py @@ -1,6 +1,8 @@ +from typing import Any + import numpy as np from scipy.integrate import odeint -from typing import List, Dict, Any + class HumanPKPDEngine: """ @@ -11,13 +13,13 @@ class HumanPKPDEngine: def __init__(self): # Default PK parameters (normalized for human adult) self.params = { - "ka": 0.5, # Absorption rate (1/hr) - "ke": 0.1, # Elimination rate (1/hr) - "k12": 0.05, # Central to peripheral 1 (1/hr) - "k21": 0.03, # Peripheral 1 to central (1/hr) - "k13": 0.02, # Central to peripheral 2 (1/hr) - "k31": 0.01, # Peripheral 2 to central (1/hr) - "Vc": 15.0, # Apparent volume of central compartment (L) + "ka": 0.5, # Absorption rate (1/hr) + "ke": 0.1, # Elimination rate (1/hr) + "k12": 0.05, # Central to peripheral 1 (1/hr) + "k21": 0.03, # Peripheral 1 to central (1/hr) + "k13": 0.02, # Central to peripheral 2 (1/hr) + "k31": 0.01, # Peripheral 2 to central (1/hr) + "Vc": 15.0, # Apparent volume of central compartment (L) } def _model_3comp(self, y, t, ka, ke, k12, k21, k13, k31, Vc): @@ -29,47 +31,41 @@ def _model_3comp(self, y, t, ka, ke, k12, k21, k13, k31, Vc): y[3]: Amount in peripheral 2 (mg) """ A_depot, A_central, A_peri1, A_peri2 = y - + # dA/dt dy0 = -ka * A_depot dy1 = ka * A_depot - (ke + k12 + k13) * A_central + k21 * A_peri1 + k31 * A_peri2 dy2 = k12 * A_central - k21 * A_peri1 dy3 = k13 * A_central - k31 * A_peri2 - + return [dy0, dy1, dy2, dy3] - def simulate_plasma_concentration( - self, - dose_mg: float, - duration_hrs: int = 48, - points: int = 100 - ) -> np.ndarray: + def simulate_plasma_concentration(self, dose_mg: float, duration_hrs: int = 48, points: int = 100) -> np.ndarray: """ Solves the 3-compartment model over the specified timeline. Returns central compartment concentration (mg/L). """ t = np.linspace(0, duration_hrs, points) y0 = [dose_mg, 0.0, 0.0, 0.0] # Initial state - + args = ( - self.params["ka"], self.params["ke"], - self.params["k12"], self.params["k21"], - self.params["k13"], self.params["k31"], - self.params["Vc"] + self.params["ka"], + self.params["ke"], + self.params["k12"], + self.params["k21"], + self.params["k13"], + self.params["k31"], + self.params["Vc"], ) - + sol = odeint(self._model_3comp, y0, t, args=args) - + # Concentration C = Amount / Volume c_central = sol[:, 1] / self.params["Vc"] return t, c_central def evaluate_pd_efficacy( - self, - concentrations: np.ndarray, - emax: float = 100.0, - ec50: float = 1.5, - gamma: float = 2.0 + self, concentrations: np.ndarray, emax: float = 100.0, ec50: float = 1.5, gamma: float = 2.0 ) -> np.ndarray: """ Sigmoidal Emax model to map concentration to therapeutic effect (0-100%). @@ -78,16 +74,16 @@ def evaluate_pd_efficacy( effect = (emax * np.power(concentrations, gamma)) / (np.power(ec50, gamma) + np.power(concentrations, gamma)) return np.nan_to_num(effect) - def run_full_simulation(self, dose_mg: float) -> Dict[str, Any]: + def run_full_simulation(self, dose_mg: float) -> dict[str, Any]: """Runs end-to-end PK/PD validation.""" t, c_central = self.simulate_plasma_concentration(dose_mg) effects = self.evaluate_pd_efficacy(c_central) - + return { "time_points": t.tolist(), "plasma_concentrations": c_central.tolist(), "therapeutic_effects": effects.tolist(), "c_max": float(np.max(c_central)), "t_max": float(t[np.argmax(c_central)]), - "auc": float(np.trapz(c_central, t)) + "auc": float(np.trapz(c_central, t)), } diff --git a/drug_discovery/testing/idiosyncratic_toxicity.py b/drug_discovery/testing/idiosyncratic_toxicity.py index a8bdfa66ea..2702a05b79 100644 --- a/drug_discovery/testing/idiosyncratic_toxicity.py +++ b/drug_discovery/testing/idiosyncratic_toxicity.py @@ -1,12 +1,15 @@ -import torch +from typing import Any + from rdkit import Chem from rdkit.Chem import Descriptors -from typing import Dict, Any + class SevereToxicityVeto(Exception): """Exception raised when a molecule fails critical human organ safety checks.""" + pass + class IdiosyncraticToxScreener: """ Identifies rare but fatal human toxicities, focusing on: @@ -36,13 +39,13 @@ def predict_dili_risk(self, smiles: str) -> float: risk_score += 0.5 if tpsa <= self.tpsa_threshold: risk_score += 0.5 - + return risk_score def flag_mitochondrial_toxicity(self, smiles: str) -> bool: """ Checks for structural alerts that decouple mitochondrial oxidative phosphorylation. - Structural alerts include: phenols with multiple halogen/nitro groups, + Structural alerts include: phenols with multiple halogen/nitro groups, lipophilic weak acids, and certain quinones. """ mol = Chem.MolFromSmiles(smiles) @@ -53,8 +56,8 @@ def flag_mitochondrial_toxicity(self, smiles: str) -> bool: alerts = { "halogenated_phenol": "Oc1c([F,Cl,Br,I])cc([F,Cl,Br,I])cc1", "nitro_phenol": "Oc1c([N+](=O)[O-])cc([N+](=O)[O-])cc1", - "lipophilic_weak_acid": "c1ccccc1-[C,N,S](=O)=O", # Generic benzoic/sulfonic acid - "quinone": "C1(=O)C=CC(=O)C=C1" + "lipophilic_weak_acid": "c1ccccc1-[C,N,S](=O)=O", # Generic benzoic/sulfonic acid + "quinone": "C1(=O)C=CC(=O)C=C1", } found_alerts = [] @@ -62,19 +65,19 @@ def flag_mitochondrial_toxicity(self, smiles: str) -> bool: pattern = Chem.MolFromSmarts(smarts) if pattern and mol.HasSubstructMatch(pattern): found_alerts.append(name) - + if found_alerts: raise SevereToxicityVeto( f"Mitochondrial toxicity alert: Found {', '.join(found_alerts)}. " "Potential oxidative phosphorylation decoupling detected." ) - + return False - def screen_molecule(self, smiles: str) -> Dict[str, Any]: + def screen_molecule(self, smiles: str) -> dict[str, Any]: """Performs safety screen and raises Veto if fatal risks are detected.""" dili_risk = self.predict_dili_risk(smiles) - + # We allow high DILI risk molecules to be flagged but not vetoed alone # unless combined with mitochondrial risk. try: @@ -88,5 +91,5 @@ def screen_molecule(self, smiles: str) -> Dict[str, Any]: "smiles": smiles, "dili_risk_score": dili_risk, "mitochondrial_safe": not mito_risk, - "organ_safety_pass": dili_risk < 1.0 and not mito_risk + "organ_safety_pass": dili_risk < 1.0 and not mito_risk, } diff --git a/drug_discovery/testing/robustness.py b/drug_discovery/testing/robustness.py index e3ac591ebf..7bf34c7790 100644 --- a/drug_discovery/testing/robustness.py +++ b/drug_discovery/testing/robustness.py @@ -96,7 +96,7 @@ def test_smiles_perturbation_robustness( self, model: Callable, smiles_list: list[str], - perturbation_types: list[str] = ["tautomer"], + perturbation_types: list[str] | None = None, tolerance: float = 0.1, ) -> dict[str, Any]: """ @@ -111,6 +111,8 @@ def test_smiles_perturbation_robustness( Returns: Dictionary with robustness metrics """ + if perturbation_types is None: + perturbation_types = ["tautomer"] results = { "total_tested": 0, "robust_predictions": 0, @@ -128,7 +130,7 @@ def test_smiles_perturbation_robustness( if isinstance(original_pred, dict): # Extract main prediction value - original_pred = list(original_pred.values())[0] + original_pred = next(iter(original_pred.values())) # Test perturbations for perturb_type in perturbation_types: @@ -139,7 +141,7 @@ def test_smiles_perturbation_robustness( perturbed_pred = model(perturbed_smiles) if isinstance(perturbed_pred, dict): - perturbed_pred = list(perturbed_pred.values())[0] + perturbed_pred = next(iter(perturbed_pred.values())) # Compute deviation deviation = abs(float(original_pred) - float(perturbed_pred)) @@ -199,8 +201,8 @@ def test_distribution_shift( # Extract prediction values if isinstance(train_preds[0], dict): - train_preds = [list(p.values())[0] for p in train_preds] - test_preds = [list(p.values())[0] for p in test_preds] + train_preds = [next(iter(p.values())) for p in train_preds] + test_preds = [next(iter(p.values())) for p in test_preds] # Compute distributional statistics results["train_pred_mean"] = np.mean(train_preds) @@ -269,7 +271,7 @@ def test_cross_validation_stability( for _, row in val_data.iterrows(): pred = model(row) if isinstance(pred, dict): - pred = list(pred.values())[0] + pred = next(iter(pred.values())) val_preds.append(pred) fold_predictions.append(val_preds) @@ -331,7 +333,7 @@ def test_adversarial_robustness( try: original_pred = model(smiles) if isinstance(original_pred, dict): - original_pred = list(original_pred.values())[0] + original_pred = next(iter(original_pred.values())) results["original_prediction"] = float(original_pred) @@ -350,7 +352,7 @@ def test_adversarial_robustness( perturbed_pred = model(perturbed_smiles) if isinstance(perturbed_pred, dict): - perturbed_pred = list(perturbed_pred.values())[0] + perturbed_pred = next(iter(perturbed_pred.values())) # Check if adversarial (prediction flipped) original_class = 1 if original_pred > 0.5 else 0 @@ -380,7 +382,7 @@ def test_environmental_robustness( ) -> dict[str, Any]: """ Test molecular robustness under environmental conditions. - + Conditions can include pH, temperature, and presence of water. """ if conditions is None: @@ -391,11 +393,7 @@ def test_environmental_robustness( {"ph": 7.4, "temp": 25.0, "desc": "Storage"}, ] - results = { - "smiles": smiles, - "condition_results": [], - "overall_stability_index": 1.0 - } + results = {"smiles": smiles, "condition_results": [], "overall_stability_index": 1.0} if Chem is None: return results @@ -408,30 +406,27 @@ def test_environmental_robustness( for cond in conditions: ph = cond.get("ph", 7.4) temp = cond.get("temp", 37.0) - + # Heuristic stability estimation # 1. pH sensitivity (hydrolysis risk) # Esters, Amides (some), Anhydrides are sensitive hydrolysis_risk = 0.0 - if mol.HasSubstructMatch(Chem.MolFromSmarts("[C;$(C=O)][O;h0][C,H]")): # Ester + if mol.HasSubstructMatch(Chem.MolFromSmarts("[C;$(C=O)][O;h0][C,H]")): # Ester hydrolysis_risk += 0.3 if abs(ph - 7.0) > 2.0 else 0.1 - if mol.HasSubstructMatch(Chem.MolFromSmarts("[C;$(C=O)][N;h0,h1][C,H]")): # Amide + if mol.HasSubstructMatch(Chem.MolFromSmarts("[C;$(C=O)][N;h0,h1][C,H]")): # Amide hydrolysis_risk += 0.1 if abs(ph - 7.0) > 4.0 else 0.05 - + # 2. Temperature sensitivity (thermal degradation) thermal_risk = 0.0 if temp > 40.0: thermal_risk = (temp - 40.0) / 100.0 - + stability = max(0.0, 1.0 - (hydrolysis_risk + thermal_risk)) stability_scores.append(stability) - - results["condition_results"].append({ - "condition": cond["desc"], - "ph": ph, - "temp": temp, - "estimated_stability": stability - }) + + results["condition_results"].append( + {"condition": cond["desc"], "ph": ph, "temp": temp, "estimated_stability": stability} + ) results["overall_stability_index"] = np.mean(stability_scores) return results @@ -472,7 +467,7 @@ def test_out_of_distribution_detection( else: # Use prediction certainty as confidence if isinstance(pred, dict): - pred = list(pred.values())[0] + pred = next(iter(pred.values())) confidence = abs(float(pred) - 0.5) * 2 in_dist_confs.append(confidence) @@ -482,7 +477,7 @@ def test_out_of_distribution_detection( out_dist_confs.append(pred["confidence"]) else: if isinstance(pred, dict): - pred = list(pred.values())[0] + pred = next(iter(pred.values())) confidence = abs(float(pred) - 0.5) * 2 out_dist_confs.append(confidence) diff --git a/drug_discovery/testing/toxicity.py b/drug_discovery/testing/toxicity.py index 48a2fef65d..ea3fbdd771 100644 --- a/drug_discovery/testing/toxicity.py +++ b/drug_discovery/testing/toxicity.py @@ -352,48 +352,46 @@ def predict_all_toxicity_endpoints( def suggest_toxicity_balancer(self, smiles: str) -> dict[str, Any]: """ Suggest a companion molecule (balancer) to mitigate toxicity. - + 50000x lesser toxic compounds by balancing toxic substances. """ tox_results = self.predict_all_toxicity_endpoints(smiles) worst_ep = tox_results["overall"]["worst_endpoint"] tox_score = tox_results["overall"]["toxicity_score"] - + if tox_score < 0.3: return {"status": "Already safe", "balancer": None} - + balancers = { "cardiotoxicity": { "name": "Dexrazoxane", "mechanism": "Iron chelation / Topoisomerase II inhibition", - "reduction_factor": 0.8 + "reduction_factor": 0.8, }, "hepatotoxicity": { "name": "N-acetylcysteine", "mechanism": "Glutathione restoration", - "reduction_factor": 0.9 + "reduction_factor": 0.9, }, "mutagenicity": { "name": "Antioxidant complex (Vitamin C/E/Alpha-lipoic acid)", "mechanism": "ROS scavenging", - "reduction_factor": 0.6 + "reduction_factor": 0.6, }, - "cytotoxicity": { - "name": "L-Carnitine", - "mechanism": "Mitochondrial support", - "reduction_factor": 0.5 - } + "cytotoxicity": {"name": "L-Carnitine", "mechanism": "Mitochondrial support", "reduction_factor": 0.5}, } - - balancer = balancers.get(worst_ep, {"name": "Generic Cytoprotectant", "mechanism": "Cellular stabilization", "reduction_factor": 0.4}) - + + balancer = balancers.get( + worst_ep, {"name": "Generic Cytoprotectant", "mechanism": "Cellular stabilization", "reduction_factor": 0.4} + ) + return { "status": "Toxic", "worst_endpoint": worst_ep, "toxicity_score": tox_score, "suggested_balancer": balancer["name"], "mechanism": balancer["mechanism"], - "estimated_balanced_toxicity": tox_score * (1.0 - balancer["reduction_factor"]) + "estimated_balanced_toxicity": tox_score * (1.0 - balancer["reduction_factor"]), } def batch_predict( diff --git a/drug_discovery/testing/uncertainty.py b/drug_discovery/testing/uncertainty.py index 0032b71f8c..ab98209948 100644 --- a/drug_discovery/testing/uncertainty.py +++ b/drug_discovery/testing/uncertainty.py @@ -220,7 +220,7 @@ def monte_carlo_dropout_uncertainty( # In practice, model would have dropout enabled during inference pred = model(input_data) if isinstance(pred, dict): - pred = list(pred.values())[0] + pred = next(iter(pred.values())) predictions.append(float(pred)) predictions = np.array(predictions) @@ -390,7 +390,7 @@ def batch_uncertainty_estimation( try: pred = model(smiles) if isinstance(pred, dict): - pred = list(pred.values())[0] + pred = next(iter(pred.values())) predictions.append(float(pred)) except Exception as e: logger.warning(f"Model prediction failed for {smiles}: {e}") diff --git a/drug_discovery/toxicity/__init__.py b/drug_discovery/toxicity/__init__.py index 2b3e63d80b..48d8422a9f 100644 --- a/drug_discovery/toxicity/__init__.py +++ b/drug_discovery/toxicity/__init__.py @@ -5,7 +5,7 @@ from .qm_mm_metabolites import ReactiveMetaboliteScreener as ReactiveMetaboliteScreener __all__ = [ - "ToxPanelScorer", "HighToxicityVeto", "ReactiveMetaboliteScreener", + "ToxPanelScorer", ] diff --git a/drug_discovery/toxicity/off_target_interactome.py b/drug_discovery/toxicity/off_target_interactome.py index 8fe52bf7fc..10d70a92b4 100644 --- a/drug_discovery/toxicity/off_target_interactome.py +++ b/drug_discovery/toxicity/off_target_interactome.py @@ -9,7 +9,6 @@ import logging import math -import os from typing import Any import torch @@ -59,11 +58,11 @@ def __init__(self, target: str, score: float, threshold: float, smiles: str) -> # Off-target definitions # --------------------------------------------------------------------------- _OFF_TARGETS: list[tuple[str, float]] = [ - ("hERG", 0.5), # Cardiac ion channel — QT prolongation risk - ("CYP3A4", 0.6), # Major CYP450 isoform — DDI / liver toxicity - ("5-HT2B", 0.55), # Serotonin receptor — valvulopathy risk - ("hNAV1.5", 0.55), # Cardiac sodium channel — arrhythmia risk - ("hKv1.5", 0.5), # Cardiac potassium channel — atrial arrhythmia + ("hERG", 0.5), # Cardiac ion channel — QT prolongation risk + ("CYP3A4", 0.6), # Major CYP450 isoform — DDI / liver toxicity + ("5-HT2B", 0.55), # Serotonin receptor — valvulopathy risk + ("hNAV1.5", 0.55), # Cardiac sodium channel — arrhythmia risk + ("hKv1.5", 0.5), # Cardiac potassium channel — atrial arrhythmia ] @@ -186,10 +185,11 @@ def __init__( self.raise_on_first = raise_on_first self.use_advanced_models = use_advanced_models self._admet_predictor = None - + if self.use_advanced_models: try: - from drug_discovery.evaluation.advanced_admet import AdvancedADMETPredictor, ADMETConfig + from drug_discovery.evaluation.advanced_admet import ADMETConfig, AdvancedADMETPredictor + self._admet_predictor = AdvancedADMETPredictor(ADMETConfig()) # In a real scenario, we'd load weights here. except ImportError: @@ -230,7 +230,7 @@ def _score_target(self, target: str, props: dict[str, float], input_mol: Any = N dummy_pos = torch.randn(3, 3) dummy_edge = torch.tensor([[0, 1, 1, 2], [1, 0, 2, 1]]) dummy_tokens = torch.randint(0, 10, (1, 10)) - + preds = self._admet_predictor(dummy_z, dummy_pos, dummy_edge, dummy_tokens) if endpoint in preds: probs = F.softmax(preds[endpoint], dim=-1) diff --git a/drug_discovery/toxicity/qm_mm_metabolites.py b/drug_discovery/toxicity/qm_mm_metabolites.py index f599f48bfb..d6745b4b89 100644 --- a/drug_discovery/toxicity/qm_mm_metabolites.py +++ b/drug_discovery/toxicity/qm_mm_metabolites.py @@ -20,6 +20,7 @@ from __future__ import annotations +import contextlib import logging import math @@ -32,7 +33,7 @@ # --------------------------------------------------------------------------- try: from rdkit import Chem # type: ignore[import-untyped] - from rdkit.Chem import AllChem, Crippen, Descriptors, rdMolDescriptors # type: ignore[import-untyped] + from rdkit.Chem import AllChem, Descriptors, rdMolDescriptors # type: ignore[import-untyped] _RDKIT = True except ImportError: # pragma: no cover @@ -81,10 +82,8 @@ def _embed_3d(mol): # type: ignore[return] result = AllChem.EmbedMolecule(mol_h, AllChem.ETKDG()) if result == -1: return None - try: + with contextlib.suppress(Exception): AllChem.MMFFOptimizeMolecule(mol_h, maxIters=200) - except Exception: - pass return mol_h @@ -118,7 +117,6 @@ def _heuristic_gap_ev(smiles: str) -> float: aromatic_rings = int(rdMolDescriptors.CalcNumAromaticRings(mol)) ring_count = int(rdMolDescriptors.CalcNumRings(mol)) tpsa = float(rdMolDescriptors.CalcTPSA(mol)) - logp = float(Crippen.MolLogP(mol)) heavy = int(mol.GetNumHeavyAtoms()) dbl_bonds = sum(1 for b in mol.GetBonds() if str(b.GetBondTypeAsDouble()) == "2.0") diff --git a/drug_discovery/training/__init__.py b/drug_discovery/training/__init__.py index 44bb8e9816..552c4fff5a 100644 --- a/drug_discovery/training/__init__.py +++ b/drug_discovery/training/__init__.py @@ -2,6 +2,7 @@ from .cryptography import EncryptionProvider as EncryptionProvider from .cryptography import PrivacyControl as PrivacyControl + __all__ = [ "EncryptionProvider", "PrivacyControl", @@ -12,7 +13,7 @@ from .federated_learning import RobustFedAvg as RobustFedAvg from .federated_node import FederatedClient as FederatedClient - __all__.extend(["FederatedServer", "RobustFedAvg", "FederatedClient"]) + __all__.extend(["FederatedClient", "FederatedServer", "RobustFedAvg"]) except ImportError: pass @@ -33,7 +34,7 @@ WarmupScheduler as WarmupScheduler, ) - __all__.extend(["AdvancedTrainer", "AdvancedTrainingConfig", "WarmupScheduler", "EMA", "EarlyStopping"]) + __all__.extend(["EMA", "AdvancedTrainer", "AdvancedTrainingConfig", "EarlyStopping", "WarmupScheduler"]) except ImportError: pass diff --git a/drug_discovery/training/active_learning_oracle.py b/drug_discovery/training/active_learning_oracle.py index fd1501bbc4..8f9bcf09b1 100644 --- a/drug_discovery/training/active_learning_oracle.py +++ b/drug_discovery/training/active_learning_oracle.py @@ -212,6 +212,7 @@ def _mol_to_pdb_block(mol_h) -> str | None: # Ray-remote ABFE function # --------------------------------------------------------------------------- if _RAY_AVAILABLE: + @ray.remote(num_cpus=1) # type: ignore[misc] def simulate_abfe_remote(smiles: str, protein_pdb: str) -> float: """Ray-distributed ABFE task. diff --git a/drug_discovery/training/advanced_training.py b/drug_discovery/training/advanced_training.py index e6947cd15f..eb1bef3815 100644 --- a/drug_discovery/training/advanced_training.py +++ b/drug_discovery/training/advanced_training.py @@ -62,7 +62,7 @@ def step(self, epoch=None, metrics=None): self.current_epoch = epoch if epoch is not None else self.current_epoch + 1 if self.current_epoch <= self.warmup_epochs: factor = self.current_epoch / max(1, self.warmup_epochs) - for pg, blr in zip(self.optimizer.param_groups, self.base_lrs): + for pg, blr in zip(self.optimizer.param_groups, self.base_lrs, strict=False): pg["lr"] = blr * factor else: if isinstance(self.base_scheduler, ReduceLROnPlateau) and metrics is not None: @@ -178,9 +178,7 @@ def fit(self, train_loader, val_loader=None): history["lr"].append(lr) history["epoch_time"].append(time.time() - t0) val_text = f"{vl:.6f}" if vl is not None else "N/A" - logger.info( - f"Epoch {epoch}/{self.config.epochs} | train={tl:.6f} | val={val_text} | lr={lr:.2e}" - ) + logger.info(f"Epoch {epoch}/{self.config.epochs} | train={tl:.6f} | val={val_text} | lr={lr:.2e}") if vl is not None and vl < best_val: best_val = vl self._save(epoch, vl, "best_model.pt") diff --git a/drug_discovery/training/closed_loop.py b/drug_discovery/training/closed_loop.py index 68f767507a..b249ccb080 100644 --- a/drug_discovery/training/closed_loop.py +++ b/drug_discovery/training/closed_loop.py @@ -104,7 +104,7 @@ def observe(self, fingerprint: np.ndarray, delta_g: float) -> None: self._best_y = delta_g def observe_batch(self, fingerprints: Sequence[np.ndarray], delta_gs: Sequence[float]) -> None: - for fp, dg in zip(fingerprints, delta_gs): + for fp, dg in zip(fingerprints, delta_gs, strict=False): self.observe(fp, dg) @property @@ -244,7 +244,7 @@ def select_top_candidates( fps = np.array([smiles_to_fingerprint(s, nbits=self.fp_dim) for s in smiles_list]) ei_values = self.expected_improvement(fps, xi=xi) - k = max(min_candidates, int(math.ceil(len(smiles_list) * top_fraction))) + k = max(min_candidates, math.ceil(len(smiles_list) * top_fraction)) k = min(k, len(smiles_list)) top_indices = np.argsort(ei_values)[-k:][::-1] return [smiles_list[i] for i in top_indices] @@ -391,9 +391,7 @@ def run_closed_loop( smiles_pool = [c["smiles"] for c in candidates] if self.surrogate.n_observations >= 2: self.surrogate.fit() - selected_smiles = self.surrogate.select_top_candidates( - smiles_pool, top_fraction=surrogate_top_fraction - ) + selected_smiles = self.surrogate.select_top_candidates(smiles_pool, top_fraction=surrogate_top_fraction) else: # Not enough data yet -- send all candidates selected_smiles = smiles_pool @@ -451,16 +449,11 @@ def run_closed_loop( # ------------------------------------------------------------------ # Oracle evaluation # ------------------------------------------------------------------ - def _evaluate_candidates_with_oracle( - self, smiles_list: list[str], target_protein: str - ) -> list[dict[str, Any]]: + def _evaluate_candidates_with_oracle(self, smiles_list: list[str], target_protein: str) -> list[dict[str, Any]]: """Run the Physics Oracle on a short-list of SMILES.""" if self.physics_oracle is None: # No oracle configured -- return mock evaluations - return [ - {"smiles": s, "delta_g": -7.0 + hash(s) % 100 / 50.0, "success": True} - for s in smiles_list - ] + return [{"smiles": s, "delta_g": -7.0 + hash(s) % 100 / 50.0, "success": True} for s in smiles_list] results = self.physics_oracle.score_batch_sync(smiles_list) return [r.as_dict() for r in results] @@ -499,9 +492,7 @@ def _generate_candidates(self, target_protein: str, num_candidates: int) -> list pool = self._SEED_SMILES for i in range(num_candidates): smiles = pool[i % len(pool)] - candidates.append( - {"id": f"iter_candidate_{i}", "smiles": smiles, "generation_method": "active_learning"} - ) + candidates.append({"id": f"iter_candidate_{i}", "smiles": smiles, "generation_method": "active_learning"}) return candidates def _evaluate_candidates(self, candidates: list[dict], target_protein: str) -> list[dict]: diff --git a/drug_discovery/training/federated_node.py b/drug_discovery/training/federated_node.py index 8d275d00a6..2adc7f2775 100644 --- a/drug_discovery/training/federated_node.py +++ b/drug_discovery/training/federated_node.py @@ -28,7 +28,7 @@ def get_parameters(self, config: dict[str, any]) -> list[torch.Tensor]: return [val.cpu().numpy() for _, val in self.model.state_dict().items()] def set_parameters(self, parameters: list[torch.Tensor]): - params_dict = zip(self.model.state_dict().keys(), parameters) + params_dict = zip(self.model.state_dict().keys(), parameters, strict=False) state_dict = OrderedDict({k: torch.tensor(v) for k, v in params_dict}) self.model.load_state_dict(state_dict, strict=True) diff --git a/drug_discovery/training/nvidia_llm_finetune.py b/drug_discovery/training/nvidia_llm_finetune.py index c328b4ab24..30d3ec9db5 100644 --- a/drug_discovery/training/nvidia_llm_finetune.py +++ b/drug_discovery/training/nvidia_llm_finetune.py @@ -2,8 +2,8 @@ from __future__ import annotations -from collections.abc import Sequence -from typing import Any, Mapping +from collections.abc import Mapping, Sequence +from typing import Any import pandas as pd @@ -88,9 +88,7 @@ def format_molecule_record(row: Mapping[str, Any] | pd.Series) -> str: if name: lines.append(f"Name: {name}") lines.append(f"SMILES: {smiles}") - lines.append( - "Purpose: Continue learning public chemistry records for local medicinal chemistry fine-tuning." - ) + lines.append("Purpose: Continue learning public chemistry records for local medicinal chemistry fine-tuning.") return "\n".join(lines) diff --git a/drug_discovery/validation/__init__.py b/drug_discovery/validation/__init__.py index 0fbafabf38..26b705f686 100644 --- a/drug_discovery/validation/__init__.py +++ b/drug_discovery/validation/__init__.py @@ -41,18 +41,18 @@ __all__.extend( [ - "set_global_seed", - "config_hash", + "CLASSIFICATION_METRICS", + "REGRESSION_METRICS", + "EliteValidationSuite", + "ExperimentReport", + "bootstrap_ci", "compute_metrics", - "scaffold_split", - "scaffold_kfold", + "config_hash", "paired_ttest", + "scaffold_kfold", + "scaffold_split", + "set_global_seed", "wilcoxon_test", - "bootstrap_ci", - "ExperimentReport", - "REGRESSION_METRICS", - "CLASSIFICATION_METRICS", - "EliteValidationSuite", ] ) except ImportError: diff --git a/drug_discovery/validation/elite_protocols.py b/drug_discovery/validation/elite_protocols.py index 98893db75e..86411e709a 100644 --- a/drug_discovery/validation/elite_protocols.py +++ b/drug_discovery/validation/elite_protocols.py @@ -344,7 +344,7 @@ def protocol_agentic_hallucination_compliance(self) -> dict[str, Any]: scores += scores[:1] angles += angles[:1] - fig, ax = plt.subplots(figsize=(8, 8), subplot_kw=dict(polar=True)) + _fig, ax = plt.subplots(figsize=(8, 8), subplot_kw=dict(polar=True)) ax.fill(angles, scores, color="teal", alpha=0.25) ax.plot(angles, scores, color="teal", linewidth=2) ax.set_theta_offset(np.pi / 2) diff --git a/drug_discovery/validation/scientific_validation.py b/drug_discovery/validation/scientific_validation.py index a13c576bd5..8522226868 100644 --- a/drug_discovery/validation/scientific_validation.py +++ b/drug_discovery/validation/scientific_validation.py @@ -156,7 +156,7 @@ def scaffold_kfold(smiles_list, n_folds=5, seed=42): smallest = min(range(n_folds), key=lambda i: sizes[i]) folds[smallest].extend(ss) sizes[smallest] += len(ss) - return [(sum([folds[j] for j in range(n_folds) if j != i], []), folds[i]) for i in range(n_folds)] + return [([item for j in range(n_folds) if j != i for item in folds[j]], folds[i]) for i in range(n_folds)] # --- Statistical tests --- diff --git a/drug_discovery/web_scraping/__init__.py b/drug_discovery/web_scraping/__init__.py index dfc33eb166..bc63eaffef 100644 --- a/drug_discovery/web_scraping/__init__.py +++ b/drug_discovery/web_scraping/__init__.py @@ -13,10 +13,10 @@ ) __all__ = [ + "AISynthesisChat", "BiomedicalScraper", - "PubMedAPI", - "WebDataProcessor", "InternetSearchClient", "OnlineResourceReader", - "AISynthesisChat", + "PubMedAPI", + "WebDataProcessor", ] diff --git a/drug_discovery/xenobiology/__init__.py b/drug_discovery/xenobiology/__init__.py index 98d443d527..507627964a 100644 --- a/drug_discovery/xenobiology/__init__.py +++ b/drug_discovery/xenobiology/__init__.py @@ -2,4 +2,4 @@ from .synthesizer import OrthogonalTranslationSimulator, XenoProteinGenerator -__all__ = ["XenoProteinGenerator", "OrthogonalTranslationSimulator"] +__all__ = ["OrthogonalTranslationSimulator", "XenoProteinGenerator"] diff --git a/drug_discovery/xenobiology/synthesizer.py b/drug_discovery/xenobiology/synthesizer.py index 3f280ccf85..2d9ef38b38 100644 --- a/drug_discovery/xenobiology/synthesizer.py +++ b/drug_discovery/xenobiology/synthesizer.py @@ -36,7 +36,9 @@ def __init__(self, synthetic_alphabet_size: int = 100): if pyrosetta is None: logger.warning("PyRosetta not installed. Xenoprotein design will use fallback modeling.") - def design_xenoprotein(self, scaffold: Pose | None = None, residues_to_mutate: list[int] = None) -> dict[str, Any]: + def design_xenoprotein( + self, scaffold: Pose | None = None, residues_to_mutate: list[int] | None = None + ) -> dict[str, Any]: """ Design a protein incorporating synthetic amino acids. """ diff --git a/examples/2024_breakthroughs.py b/examples/2024_breakthroughs.py index fc4d17b318..014e364428 100644 --- a/examples/2024_breakthroughs.py +++ b/examples/2024_breakthroughs.py @@ -3,10 +3,13 @@ """ import asyncio + from drug_discovery.alphafold3.alphafold3_docking import AlphaFold3Docking from drug_discovery.rfdiffusion.protein_design import RFDiffusionDesigner + # etc. + async def main(): # AF3 proxy af3 = AlphaFold3Docking("protein_seq") @@ -20,15 +23,18 @@ async def main(): # CRISPR base edit from models.biologics.crispr_base_editor import CRISPRBaseEditor + editor = CRISPRBaseEditor() edit = editor.base_edit("ATCG", 1, "A", "G") print(edit) # ADC from models.nextgen_adcs.adc_optimizer import ADCOptimizer + adc = ADCOptimizer() res = adc.optimize("DM1", "Her2-shuttle") print(res) + if __name__ == "__main__": - asyncio.run(main()) \ No newline at end of file + asyncio.run(main()) diff --git a/examples/admet_prediction.py b/examples/admet_prediction.py index 450423a14f..d26c368a9d 100644 --- a/examples/admet_prediction.py +++ b/examples/admet_prediction.py @@ -2,9 +2,10 @@ Example: ADMET Property Prediction """ -from drug_discovery.evaluation import ADMETPredictor import pandas as pd +from drug_discovery.evaluation import ADMETPredictor + def main(): # Initialize ADMET predictor @@ -12,10 +13,10 @@ def main(): # Example molecules molecules = { - 'Aspirin': 'CC(=O)OC1=CC=CC=C1C(=O)O', - 'Ibuprofen': 'CC(C)CC1=CC=C(C=C1)C(C)C(=O)O', - 'Caffeine': 'CN1C=NC2=C1C(=O)N(C(=O)N2C)C', - 'Penicillin': 'CC1(C)SC2C(NC(=O)Cc3ccccc3)C(=O)N2C1C(=O)O', + "Aspirin": "CC(=O)OC1=CC=CC=C1C(=O)O", + "Ibuprofen": "CC(C)CC1=CC=C(C=C1)C(C)C(=O)O", + "Caffeine": "CN1C=NC2=C1C(=O)N(C(=O)N2C)C", + "Penicillin": "CC1(C)SC2C(NC(=O)Cc3ccccc3)C(=O)N2C1C(=O)O", } results = [] @@ -32,10 +33,10 @@ def main(): print("\nLipinski's Rule of Five:\n Could not compute properties") print("\n" + "=" * 50 + "\n") continue - print(f"\nLipinski's Rule of Five:") + print("\nLipinski's Rule of Five:") print(f" Pass: {lipinski['passes']}") print(f" Violations: {lipinski['num_violations']}") - if lipinski['violations']: + if lipinski["violations"]: print(f" Issues: {', '.join(lipinski['violations'])}") # Drug-likeness @@ -58,30 +59,32 @@ def main(): toxicity = admet.predict_toxicity_flags(smiles) if toxicity is None: toxicity = {} - print(f"\nToxicity Flags:") + print("\nToxicity Flags:") for flag, value in toxicity.items(): print(f" {flag}: {'Yes' if value else 'No'}") # Calculate properties - props = lipinski['properties'] - print(f"\nMolecular Properties:") + props = lipinski["properties"] + print("\nMolecular Properties:") print(f" MW: {props['molecular_weight']:.2f}") print(f" LogP: {props['logp']:.2f}") print(f" H-bond donors: {props['h_bond_donors']}") print(f" H-bond acceptors: {props['h_bond_acceptors']}") - print("\n" + "="*50 + "\n") + print("\n" + "=" * 50 + "\n") # Store results - results.append({ - 'name': name, - 'smiles': smiles, - 'lipinski_pass': lipinski['passes'], - 'qed': qed, - 'sa_score': sa_score, - 'mw': props['molecular_weight'], - 'logp': props['logp'], - }) + results.append( + { + "name": name, + "smiles": smiles, + "lipinski_pass": lipinski["passes"], + "qed": qed, + "sa_score": sa_score, + "mw": props["molecular_weight"], + "logp": props["logp"], + } + ) # Create summary DataFrame df = pd.DataFrame(results) @@ -89,9 +92,9 @@ def main(): print(df.to_string(index=False)) # Filter drug-like molecules - drug_like = df[(df['lipinski_pass'] == True) & (df['qed'] > 0.5)] + drug_like = df[(df["lipinski_pass"]) & (df["qed"] > 0.5)] print(f"\n✓ Drug-like molecules: {len(drug_like)}/{len(df)}") - print(drug_like['name'].tolist()) + print(drug_like["name"].tolist()) if __name__ == "__main__": diff --git a/examples/basic_usage.py b/examples/basic_usage.py index 805ae6212a..c94900a13a 100644 --- a/examples/basic_usage.py +++ b/examples/basic_usage.py @@ -4,19 +4,18 @@ from drug_discovery import DrugDiscoveryPipeline + def main(): # Initialize pipeline print("Initializing Drug Discovery Pipeline...") pipeline = DrugDiscoveryPipeline( - model_type='gnn', # Options: 'gnn', 'transformer', 'ensemble' - device='cuda' # or 'cpu' + model_type="gnn", device="cuda" # Options: 'gnn', 'transformer', 'ensemble' # or 'cpu' ) # Step 1: Collect data from public sources print("\nStep 1: Collecting molecular data...") data = pipeline.collect_data( - sources=['pubchem', 'chembl', 'approved_drugs'], - limit_per_source=500 # Start small for demo + sources=["pubchem", "chembl", "approved_drugs"], limit_per_source=500 # Start small for demo ) print(f"Collected {len(data)} molecules") @@ -25,20 +24,16 @@ def main(): # Step 2: Prepare datasets print("\nStep 2: Preparing datasets...") train_loader, test_loader = pipeline.prepare_datasets( - data=data, - smiles_col='smiles', - target_col=None, # Unsupervised for now - test_size=0.2, - batch_size=32 + data=data, smiles_col="smiles", target_col=None, test_size=0.2, batch_size=32 # Unsupervised for now ) # Step 3: Build and train model print("\nStep 3: Training model...") - history = pipeline.train( + pipeline.train( train_loader=train_loader, val_loader=test_loader, num_epochs=10, # Use more epochs in production - learning_rate=1e-4 + learning_rate=1e-4, ) # Step 4: Predict properties for a molecule @@ -52,21 +47,18 @@ def main(): # Step 5: Generate drug candidates print("\nStep 5: Generating drug candidates...") - candidates = pipeline.generate_candidates( - target_protein="EGFR", - num_candidates=10 - ) + candidates = pipeline.generate_candidates(target_protein="EGFR", num_candidates=10) print("\nTop drug candidates:") - print(candidates[['smiles', 'qed_score', 'lipinski_pass']].head()) + print(candidates[["smiles", "qed_score", "lipinski_pass"]].head()) # Step 6: Evaluate model print("\nStep 6: Evaluating model...") - metrics = pipeline.evaluate(test_loader) + pipeline.evaluate(test_loader) # Save pipeline print("\nSaving pipeline...") - pipeline.save('./checkpoints/pipeline.pt') + pipeline.save("./checkpoints/pipeline.pt") print("\n✓ Pipeline demonstration complete!") diff --git a/examples/continuous_learning.py b/examples/continuous_learning.py index b61559a1e2..f476d4bdbf 100644 --- a/examples/continuous_learning.py +++ b/examples/continuous_learning.py @@ -2,13 +2,15 @@ Example: Advanced Usage - Continuous Learning """ +import time + from drug_discovery import DrugDiscoveryPipeline from drug_discovery.training import ContinuousLearner -import time + def main(): # Initialize pipeline - pipeline = DrugDiscoveryPipeline(model_type='gnn') + pipeline = DrugDiscoveryPipeline(model_type="gnn") # Initial training print("Initial training phase...") @@ -23,7 +25,7 @@ def main(): continuous_learner = ContinuousLearner( trainer=pipeline.trainer, data_collector=pipeline.data_collector, - retrain_threshold=500 # Retrain after 500 new samples + retrain_threshold=500, # Retrain after 500 new samples ) # Simulate continuous learning loop @@ -51,10 +53,10 @@ def main(): data = combined_data # Make predictions on new molecules - if not new_data.empty and 'smiles' in new_data.columns: - sample_smiles = new_data['smiles'].iloc[0] + if not new_data.empty and "smiles" in new_data.columns: + sample_smiles = new_data["smiles"].iloc[0] properties = pipeline.predict_properties(sample_smiles) - print(f"\nPredictions for new molecule:") + print("\nPredictions for new molecule:") print(f" SMILES: {sample_smiles}") print(f" QED: {properties.get('qed_score', 'N/A')}") diff --git a/examples/simulated_combination_screen.py b/examples/simulated_combination_screen.py index 802902952f..f5ac8b2433 100644 --- a/examples/simulated_combination_screen.py +++ b/examples/simulated_combination_screen.py @@ -11,10 +11,10 @@ from __future__ import annotations +import sys from dataclasses import dataclass from itertools import combinations from pathlib import Path -import sys import pandas as pd diff --git a/external/supply_chain.py b/external/supply_chain.py index bfc61410b3..956dc332e8 100644 --- a/external/supply_chain.py +++ b/external/supply_chain.py @@ -17,9 +17,9 @@ import hashlib import logging -import time +from collections.abc import Sequence from dataclasses import dataclass, field -from typing import Any, Sequence +from typing import Any logger = logging.getLogger(__name__) @@ -258,7 +258,7 @@ def _brics_decompose(smiles: str) -> list[str]: clean = Chem.MolFromSmiles(frag) if clean is not None: # Replace dummy atoms with H - from rdkit.Chem import AllChem, RWMol + from rdkit.Chem import RWMol rw = RWMol(clean) atoms_to_remove = [] @@ -381,13 +381,11 @@ def _check_availability(self, fragment_smiles: str) -> Synthon: if result.get("catalog_id"): synthon.catalog_ids[vendor.name] = result["catalog_id"] price = result.get("price_usd") - if price is not None: - if best_price is None or price < best_price: - best_price = price + if price is not None and (best_price is None or price < best_price): + best_price = price lead = result.get("lead_time_days") - if lead is not None: - if best_lead is None or lead < best_lead: - best_lead = lead + if lead is not None and (best_lead is None or lead < best_lead): + best_lead = lead except Exception as exc: logger.warning("Vendor %s query failed for %s: %s", vendor.name, fragment_smiles, exc) diff --git a/external_plugins/vaccine_trigger_cli.py b/external_plugins/vaccine_trigger_cli.py index 7e465b415b..ef467452fe 100644 --- a/external_plugins/vaccine_trigger_cli.py +++ b/external_plugins/vaccine_trigger_cli.py @@ -1,28 +1,28 @@ import argparse -import importlib import json import logging -import sys import os +import sys # Configure logging -logging.basicConfig(level=logging.INFO, format='%(name)s - %(levelname)s - %(message)s') +logging.basicConfig(level=logging.INFO, format="%(name)s - %(levelname)s - %(message)s") logger = logging.getLogger("VaccineTrigger") + def run_generate_mrna(args): """ On-demand orchestration for mRNA vaccine design. Loads heavy modules only when needed. """ logger.info(f"Triggering mRNA generation for {args.viral_fasta}...") - + # Dynamic imports to save memory/VRAM try: from external_plugins.vaccinology.mhc_epitope_mapper import PatientEpitopeMapper from external_plugins.vaccinology.mrna_compiler import ThermodynamicmRNACompiler from external_plugins.vaccinology.prefusion_stabilizer import PrefusionLockEngine except ImportError as e: - logger.error(f"Failed to load heavy vaccinology modules: {str(e)}") + logger.error(f"Failed to load heavy vaccinology modules: {e!s}") sys.exit(1) # 1. Epitope Mapping @@ -30,15 +30,15 @@ def run_generate_mrna(args): # Mock HLA alleles if not provided hla_alleles = ["HLA-A*02:01"] if args.patient_hla and os.path.exists(args.patient_hla): - with open(args.patient_hla, 'r') as f: + with open(args.patient_hla) as f: hla_alleles = json.load(f).get("alleles", hla_alleles) - + antigen_seq = mapper.select_optimal_antigen_payload(args.viral_fasta, hla_alleles) - + # 2. mRNA Compilation compiler = ThermodynamicmRNACompiler() payload = compiler.compile_full_payload(antigen_seq) - + # 3. Structural Stabilization (if PDB is provided) stability_report = {} if args.viral_pdb: @@ -48,7 +48,7 @@ def run_generate_mrna(args): stability_report = { "stabilizing_mutations": mutations, "predicted_ddg": ddg, - "conformation": "prefusion_locked" + "conformation": "prefusion_locked", } # Final Package @@ -57,15 +57,16 @@ def run_generate_mrna(args): "target_fasta": args.viral_fasta, "optimized_payload": payload, "structural_validation": stability_report, - "patient_context": {"hla_alleles": hla_alleles} + "patient_context": {"hla_alleles": hla_alleles}, } - + output_file = args.output or "vaccine_blueprint.json" - with open(output_file, 'w') as f: + with open(output_file, "w") as f: json.dump(blueprint, f, indent=4) - + logger.info(f"Vaccine design complete. Blueprint saved to {output_file}") + def main(): parser = argparse.ArgumentParser(description="ZANE External Vaccine Generation CLI") subparsers = parser.add_subparsers(dest="command", help="Vaccine commands") @@ -84,5 +85,6 @@ def main(): else: parser.print_help() + if __name__ == "__main__": main() diff --git a/external_plugins/vaccinology/mhc_epitope_mapper.py b/external_plugins/vaccinology/mhc_epitope_mapper.py index d0445d059c..8a8a6c5f4c 100644 --- a/external_plugins/vaccinology/mhc_epitope_mapper.py +++ b/external_plugins/vaccinology/mhc_epitope_mapper.py @@ -1,23 +1,22 @@ +import logging + import pandas as pd import torch import torch.nn as nn from Bio import SeqIO -from typing import List, Dict, Any, Optional -import logging logger = logging.getLogger(__name__) + class PatientEpitopeMapper: """ Simulates binding affinity of viral peptides against patient-specific MHC alleles. """ + def __init__(self): # Mock neural network for binding prediction (similar to NetMHCpan architecture) self.model = nn.Sequential( - nn.Linear(20 * 9, 128), # 9-mer peptides, 20 amino acids - nn.ReLU(), - nn.Linear(128, 1), - nn.Sigmoid() + nn.Linear(20 * 9, 128), nn.ReLU(), nn.Linear(128, 1), nn.Sigmoid() # 9-mer peptides, 20 amino acids ) self.aa_map = {aa: i for i, aa in enumerate("ACDEFGHIKLMNPQRSTVWY")} @@ -29,7 +28,7 @@ def _encode_peptide(self, peptide: str) -> torch.Tensor: tensor[i, self.aa_map[aa]] = 1.0 return tensor.flatten() - def predict_mhc_class_I_binding(self, viral_fasta: str, patient_hla_alleles: List[str]) -> pd.DataFrame: + def predict_mhc_class_I_binding(self, viral_fasta: str, patient_hla_alleles: list[str]) -> pd.DataFrame: """ Predicts binding affinity (K_d) for all 9-mer peptides in the viral sequence. """ @@ -40,47 +39,49 @@ def predict_mhc_class_I_binding(self, viral_fasta: str, patient_hla_alleles: Lis sequence = str(record.seq) # Sliding window of 9 amino acids for i in range(len(sequence) - 8): - peptide = sequence[i:i+9] - + peptide = sequence[i : i + 9] + # Simulated prediction logic # In a real system, this would be conditioned on the HLA allele peptide_tensor = self._encode_peptide(peptide) binding_score = self.model(peptide_tensor).item() - + # Convert score to nanomolar Kd (Log scale mapping) # High score = Low Kd (strong binding) - kd_nm = 50000**(1 - binding_score) - - results.append({ - "peptide": peptide, - "start_pos": i, - "kd_nm": kd_nm, - "allele": patient_hla_alleles[0] if patient_hla_alleles else "HLA-A*02:01", - "immunogenicity_score": binding_score - }) - + kd_nm = 50000 ** (1 - binding_score) + + results.append( + { + "peptide": peptide, + "start_pos": i, + "kd_nm": kd_nm, + "allele": patient_hla_alleles[0] if patient_hla_alleles else "HLA-A*02:01", + "immunogenicity_score": binding_score, + } + ) + df = pd.DataFrame(results) # Filter for "strong binders" (Kd < 50nM) return df.sort_values(by="kd_nm").reset_index(drop=True) - + except Exception as e: - logger.error(f"Epitope mapping failed: {str(e)}") + logger.error(f"Epitope mapping failed: {e!s}") return pd.DataFrame() - def select_optimal_antigen_payload(self, viral_fasta: str, patient_hla_alleles: List[str]) -> str: + def select_optimal_antigen_payload(self, viral_fasta: str, patient_hla_alleles: list[str]) -> str: """ Constructs a synthetic poly-epitope chain containing highly immunogenic fragments. """ predictions = self.predict_mhc_class_I_binding(viral_fasta, patient_hla_alleles) - + if predictions.empty: return "" - + # Select top 5 non-overlapping epitopes top_epitopes = predictions.head(5)["peptide"].tolist() - + # Link epitopes with flexible spacers (e.g., AAY linkers) poly_epitope = "AAY".join(top_epitopes) - + logger.info(f"Generated optimized poly-epitope payload of length {len(poly_epitope)}") return poly_epitope diff --git a/external_plugins/vaccinology/mrna_compiler.py b/external_plugins/vaccinology/mrna_compiler.py index 2156fc14f7..e3de4b1e27 100644 --- a/external_plugins/vaccinology/mrna_compiler.py +++ b/external_plugins/vaccinology/mrna_compiler.py @@ -1,26 +1,25 @@ -import networkx as nx -from typing import Dict, List, Optional import logging -import random logger = logging.getLogger(__name__) + class ThermodynamicmRNACompiler: """ Optimizes mRNA sequences for maximum protein expression and thermodynamic stability. """ + def __init__(self): # Standard human codon usage bias (simplified) self.codon_table = { - 'A': ['GCC', 'GCT', 'GCA', 'GCG'], - 'L': ['CTG', 'CTC', 'TTG', 'TTA', 'CTT', 'CTA'], - 'P': ['CCC', 'CCT', 'CCA', 'CCG'], - 'R': ['CGC', 'AGG', 'CGT', 'AGA', 'CGA', 'CGG'], - 'V': ['GTG', 'GTC', 'GTT', 'GTA'], + "A": ["GCC", "GCT", "GCA", "GCG"], + "L": ["CTG", "CTC", "TTG", "TTA", "CTT", "CTA"], + "P": ["CCC", "CCT", "CCA", "CCG"], + "R": ["CGC", "AGG", "CGT", "AGA", "CGA", "CGG"], + "V": ["GTG", "GTC", "GTT", "GTA"], # ... (mapping would be complete in production) } # Preferred codons for humans - self.preferred_codons = {'A': 'GCC', 'L': 'CTG', 'P': 'CCC', 'R': 'CGC', 'V': 'GTG'} + self.preferred_codons = {"A": "GCC", "L": "CTG", "P": "CCC", "R": "CGC", "V": "GTG"} def optimize_codon_adaptation_index(self, amino_acid_seq: str) -> str: """ @@ -32,19 +31,19 @@ def optimize_codon_adaptation_index(self, amino_acid_seq: str) -> str: mrna_seq.append(self.preferred_codons[aa]) else: # Fallback to random if not in preferred (simulated) - mrna_seq.append('AUG') # Simplified fallback - + mrna_seq.append("AUG") # Simplified fallback + optimized_seq = "".join(mrna_seq) logger.info(f"Codon optimization complete. Sequence length: {len(optimized_seq)}nt") return optimized_seq def minimize_mrna_free_energy(self, rna_sequence: str) -> str: """ - Iteratively tweaks synonymous codons to achieve a deeply negative + Iteratively tweaks synonymous codons to achieve a deeply negative Minimum Free Energy (MFE) using structural prediction. """ current_seq = rna_sequence - + # Lazy import of ViennaRNA/RNA try: import RNA @@ -54,22 +53,22 @@ def minimize_mrna_free_energy(self, rna_sequence: str) -> str: def get_mfe(seq): if RNA: - (ss, mfe) = RNA.fold(seq) + _ss, mfe = RNA.fold(seq) return mfe else: # Mock MFE: High GC content generally lowers MFE - gc_content = (seq.count('G') + seq.count('C')) / len(seq) - return -100.0 * gc_content # Simulated MFE + gc_content = (seq.count("G") + seq.count("C")) / len(seq) + return -100.0 * gc_content # Simulated MFE best_mfe = get_mfe(current_seq) - + # Iterative optimization (Monte Carlo approach) - for i in range(10): # Simplified 10 iterations + for _i in range(10): # Simplified 10 iterations # Propose a synonymous codon swap # (Logic omitted for brevity - would maintain amino acid identity) - test_seq = current_seq # Simulate a swap + test_seq = current_seq # Simulate a swap test_mfe = get_mfe(test_seq) - + if test_mfe < best_mfe: best_mfe = test_mfe current_seq = test_seq @@ -77,17 +76,17 @@ def get_mfe(seq): logger.info(f"Thermodynamic optimization complete. Final MFE: {best_mfe:.2f} kcal/mol") return current_seq - def compile_full_payload(self, protein_seq: str) -> Dict[str, str]: + def compile_full_payload(self, protein_seq: str) -> dict[str, str]: """ - Executes full mRNA compilation pipeline: UTR addition, 1-methylpseudouridine + Executes full mRNA compilation pipeline: UTR addition, 1-methylpseudouridine incorporation (simulated), and thermodynamic folding. """ cai_optimized = self.optimize_codon_adaptation_index(protein_seq) stable_mrna = self.minimize_mrna_free_energy(cai_optimized) - + return { "mrna_sequence": stable_mrna, "modifications": "1-methylpseudouridine", "5_utr": "GGGAUAAUACUCAUACUAUUCCCGAGUAUUACUAUACUCCCAUCG", - "3_utr": "UUUGAAUUAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA" + "3_utr": "UUUGAAUUAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA", } diff --git a/external_plugins/vaccinology/prefusion_stabilizer.py b/external_plugins/vaccinology/prefusion_stabilizer.py index e70f71ed1c..d517cb4783 100644 --- a/external_plugins/vaccinology/prefusion_stabilizer.py +++ b/external_plugins/vaccinology/prefusion_stabilizer.py @@ -1,17 +1,19 @@ -import numpy as np -from Bio.PDB import PDBParser, Selection -from typing import List, Dict, Tuple, Optional import logging +import numpy as np +from Bio.PDB import PDBParser + # Set up logging logger = logging.getLogger(__name__) + class PrefusionLockEngine: """ Engine for stabilizing viral fusion proteins in their prefusion conformation to enhance vaccine efficacy. """ - def __init__(self, pdb_path: Optional[str] = None): + + def __init__(self, pdb_path: str | None = None): self.pdb_path = pdb_path self.parser = PDBParser(QUIET=True) self.rosetta_initialized = False @@ -21,80 +23,81 @@ def _init_pyrosetta(self): if not self.rosetta_initialized: try: import pyrosetta + pyrosetta.init(extra_options="-constant_seed -mute all") self.rosetta_initialized = True logger.info("PyRosetta initialized successfully.") except ImportError: logger.warning("PyRosetta not found. Using structural mocking for delta-delta G calculations.") - def calculate_mutation_ddg(self, pdb_path: str, mutation_list: List[str]) -> float: + def calculate_mutation_ddg(self, pdb_path: str, mutation_list: list[str]) -> float: """ Calculates the change in folding free energy (ddG) for a set of mutations. A negative value indicates stabilization. """ self._init_pyrosetta() - + if self.rosetta_initialized: import pyrosetta from pyrosetta.rosetta.core.scoring import get_score_function - + pose = pyrosetta.pose_from_pdb(pdb_path) scorefxn = get_score_function() - + # Baseline score initial_score = scorefxn(pose) - + # Apply mutations (simplified logic) for mut in mutation_list: # Expecting format like 'A123P' (Wildtype-ResidueIndex-Mutation) res_idx = int(mut[1:-1]) new_aa = mut[-1] pyrosetta.toolbox.mutants.mutate_residue(pose, res_idx, new_aa) - + final_score = scorefxn(pose) return float(final_score - initial_score) else: # Structural mocking: Proline substitutions in loops typically stabilize ddg = 0.0 for mut in mutation_list: - if mut.endswith('P'): - ddg -= 1.5 # Heuristic stabilization for Proline - if mut.startswith('C') and 'C' in mut[1:]: # Disulfide bond + if mut.endswith("P"): + ddg -= 1.5 # Heuristic stabilization for Proline + if mut.startswith("C") and "C" in mut[1:]: # Disulfide bond ddg -= 3.0 return ddg - def scan_for_metastable_hinges(self, pdb_path: str) -> List[int]: + def scan_for_metastable_hinges(self, pdb_path: str) -> list[int]: """ Identifies flexible hinge regions based on B-factors or local geometry. These are targets for stabilizing mutations. """ structure = self.parser.get_structure("protein", pdb_path) model = structure[0] - + hinge_indices = [] residues = list(model.get_residues()) - + # Calculate local displacement or B-factor variance for i in range(1, len(residues) - 1): res = residues[i] # Simple heuristic: high B-factors often correlate with flexibility avg_b = np.mean([atom.get_bfactor() for atom in res.get_atoms()]) - - if avg_b > 50.0: # Threshold for high flexibility + + if avg_b > 50.0: # Threshold for high flexibility hinge_indices.append(res.get_id()[1]) - + logger.info(f"Identified {len(hinge_indices)} metastable hinge candidates in {pdb_path}") return hinge_indices - def suggest_stabilizing_mutations(self, pdb_path: str) -> List[str]: + def suggest_stabilizing_mutations(self, pdb_path: str) -> list[str]: """ Autonomously suggests mutations (e.g., 2P stabilization) to freeze the geometry. """ hinges = self.scan_for_metastable_hinges(pdb_path) suggestions = [] - + # Focus on top 2 most flexible regions for Proline substitution for idx in hinges[:2]: - suggestions.append(f"X{idx}P") # X represents original AA - + suggestions.append(f"X{idx}P") # X represents original AA + return suggestions diff --git a/infrastructure/api/batch_celery_processor.py b/infrastructure/api/batch_celery_processor.py index de44b53939..eefad8d4e9 100644 --- a/infrastructure/api/batch_celery_processor.py +++ b/infrastructure/api/batch_celery_processor.py @@ -1,8 +1,10 @@ +import logging import os +from typing import Any + from celery import Celery -from typing import List, Dict, Any -import logging -from drug_discovery.evaluation.advanced_admet import AdvancedADMETPredictor, ADMETConfig + +from drug_discovery.evaluation.advanced_admet import ADMETConfig, AdvancedADMETPredictor # Initialize Celery with Redis backend REDIS_URL = os.getenv("REDIS_URL", "redis://localhost:6379/0") @@ -10,6 +12,7 @@ logger = logging.getLogger(__name__) + class HighThroughputBatchQueue: """ Orchestrates high-volume proprietary molecule scoring across EKS GPU workers. @@ -17,62 +20,52 @@ class HighThroughputBatchQueue: @staticmethod @celery_app.task(bind=True, name="process_pharma_library_batch") - def process_pharma_library_batch(self, batch_id: str, smiles_list: List[str]) -> Dict[str, Any]: + def process_pharma_library_batch(self, batch_id: str, smiles_list: list[str]) -> dict[str, Any]: """ Runs full ADMET and ABFE scoring pipeline for a library batch. """ total = len(smiles_list) results = [] - + # Load Predictor predictor = AdvancedADMETPredictor(ADMETConfig()) - + for i, smiles in enumerate(smiles_list): try: # Perform scoring # Note: In production, we'd use the parallel/Ray logic implemented previously - score = predictor.forward_smiles(smiles) # Simplified call + score = predictor.forward_smiles(smiles) # Simplified call results.append({"smiles": smiles, "score": score}) - + # Update progress - self.update_state(state='PROGRESS', meta={'current': i, 'total': total, 'percent': (i / total) * 100}) + self.update_state(state="PROGRESS", meta={"current": i, "total": total, "percent": (i / total) * 100}) except Exception as e: logger.error(f"Error scoring {smiles}: {e}") results.append({"smiles": smiles, "error": str(e)}) - return { - "batch_id": batch_id, - "status": "COMPLETED", - "total_processed": total, - "results": results - } + return {"batch_id": batch_id, "status": "COMPLETED", "total_processed": total, "results": results} @staticmethod - def get_task_status(task_id: str) -> Dict[str, Any]: + def get_task_status(task_id: str) -> dict[str, Any]: """Allows polling of completion percentage.""" task = celery_app.AsyncResult(task_id) - if task.state == 'PENDING': - response = { - 'state': task.state, - 'current': 0, - 'total': 1, - 'status': 'Pending...' - } - elif task.state != 'FAILURE': + if task.state == "PENDING": + response = {"state": task.state, "current": 0, "total": 1, "status": "Pending..."} + elif task.state != "FAILURE": response = { - 'state': task.state, - 'current': task.info.get('current', 0) if isinstance(task.info, dict) else 0, - 'total': task.info.get('total', 1) if isinstance(task.info, dict) else 1, - 'percent': task.info.get('percent', 0) if isinstance(task.info, dict) else 0, - 'status': task.info.get('status', '') if isinstance(task.info, dict) else '' + "state": task.state, + "current": task.info.get("current", 0) if isinstance(task.info, dict) else 0, + "total": task.info.get("total", 1) if isinstance(task.info, dict) else 1, + "percent": task.info.get("percent", 0) if isinstance(task.info, dict) else 0, + "status": task.info.get("status", "") if isinstance(task.info, dict) else "", } - if 'results' in task.info if isinstance(task.info, dict) else False: - response['results'] = task.info['results'] + if "results" in task.info if isinstance(task.info, dict) else False: + response["results"] = task.info["results"] else: response = { - 'state': task.state, - 'current': 1, - 'total': 1, - 'status': str(task.info), + "state": task.state, + "current": 1, + "total": 1, + "status": str(task.info), } return response diff --git a/infrastructure/api/enterprise_sso.py b/infrastructure/api/enterprise_sso.py index 2dc88769cc..75fd410ae1 100644 --- a/infrastructure/api/enterprise_sso.py +++ b/infrastructure/api/enterprise_sso.py @@ -1,10 +1,10 @@ +import os + from fastapi import Depends, HTTPException, status from fastapi.security import OAuth2PasswordBearer from jose import JWTError, jwt from passlib.context import CryptContext from pydantic import BaseModel -from typing import List, Optional -import os # Security Config SECRET_KEY = os.getenv("JWT_SECRET_KEY", "SUPER_SECRET_PHARMA_KEY") @@ -14,13 +14,16 @@ pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token") + class User(BaseModel): username: str - role: str # Junior_Chemist, Senior_Scientist, Lab_Director + role: str # Junior_Chemist, Senior_Scientist, Lab_Director + class TokenData(BaseModel): - username: Optional[str] = None - role: Optional[str] = None + username: str | None = None + role: str | None = None + class EnterpriseSSO: """ @@ -29,8 +32,9 @@ class EnterpriseSSO: """ @staticmethod - def verify_role(required_roles: List[str]): + def verify_role(required_roles: list[str]): """RBAC Dependency: Ensures the user has the required seniority.""" + async def role_checker(token: str = Depends(oauth2_scheme)): credentials_exception = HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -50,10 +54,10 @@ async def role_checker(token: str = Depends(oauth2_scheme)): if token_data.role not in required_roles: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, - detail=f"Access denied. Required roles: {required_roles}. Current role: {token_data.role}" + detail=f"Access denied. Required roles: {required_roles}. Current role: {token_data.role}", ) return token_data - + return role_checker @staticmethod @@ -65,6 +69,7 @@ def get_current_user(token: str = Depends(oauth2_scheme)): except JWTError: raise HTTPException(status_code=401, detail="Invalid token") + # Usage Examples for Endpoints: # @app.post("/export/dataset", dependencies=[Depends(EnterpriseSSO.verify_role(["Lab_Director"]))]) # async def export_dataset(): diff --git a/infrastructure/api_gateway.py b/infrastructure/api_gateway.py index 77d225b7e3..0a30cbd644 100644 --- a/infrastructure/api_gateway.py +++ b/infrastructure/api_gateway.py @@ -19,11 +19,9 @@ from __future__ import annotations -import hashlib import logging import time import uuid -from dataclasses import dataclass, field from enum import Enum from typing import Any @@ -54,7 +52,7 @@ def Field(default: Any = None, **kw: Any) -> Any: # type: ignore[misc] try: - from fastapi import FastAPI, HTTPException, BackgroundTasks, Depends + from fastapi import BackgroundTasks, Depends, FastAPI, HTTPException _FASTAPI = True except ImportError: @@ -191,7 +189,7 @@ def _execute_generation_task(task_id: str, request_data: dict[str, Any]) -> dict t0 = time.monotonic() try: - from drug_discovery.safety.end_to_end_pipeline import SafeGenerationPipeline, PipelineConfig + from drug_discovery.safety.end_to_end_pipeline import PipelineConfig, SafeGenerationPipeline cfg = PipelineConfig( num_candidates=request_data.get("num_candidates", 100), @@ -206,14 +204,16 @@ def _execute_generation_task(task_id: str, request_data: dict[str, Any]) -> dict candidates = [] for c in result.final_candidates: scores = c.get("scores", {}) - candidates.append({ - "smiles": c.get("smiles", ""), - "delta_g": scores.get("delta_g"), - "toxicity": scores.get("toxicity", 0.0), - "drug_likeness": scores.get("drug_likeness", 0.0), - "sa_score": scores.get("sa_score", 3.0), - "pareto_rank": c.get("pareto_rank", 0), - }) + candidates.append( + { + "smiles": c.get("smiles", ""), + "delta_g": scores.get("delta_g"), + "toxicity": scores.get("toxicity", 0.0), + "drug_likeness": scores.get("drug_likeness", 0.0), + "sa_score": scores.get("sa_score", 3.0), + "pareto_rank": c.get("pareto_rank", 0), + } + ) task_result = { "task_id": task_id, diff --git a/infrastructure/cloud_lab/__init__.py b/infrastructure/cloud_lab/__init__.py index 6a4ee1c6e9..13cc8f0c99 100644 --- a/infrastructure/cloud_lab/__init__.py +++ b/infrastructure/cloud_lab/__init__.py @@ -1,3 +1,3 @@ from .os_kernel import OSKernel -__all__ = ["OSKernel"] \ No newline at end of file +__all__ = ["OSKernel"] diff --git a/infrastructure/cloud_lab/os_kernel.py b/infrastructure/cloud_lab/os_kernel.py index b0bfb6a144..bda5b91e82 100644 --- a/infrastructure/cloud_lab/os_kernel.py +++ b/infrastructure/cloud_lab/os_kernel.py @@ -9,7 +9,12 @@ def submit(self, spec: dict) -> str: def wait(self, job_id: str) -> dict: print(f"Mock Bacalhau wait for {job_id}") - return {"status": "success", "results": ["/ipfs/mock_crispr_data.json"], "outputs": {"data": "optimized results"}} + return { + "status": "success", + "results": ["/ipfs/mock_crispr_data.json"], + "outputs": {"data": "optimized results"}, + } + class OSKernel: def __init__(self): @@ -23,24 +28,20 @@ def compile_labop(self, protocol_str: str) -> dict[str, Any]: try: # LabOP integration stub from labop.core import Protocol # assume pip install labop + protocol = Protocol.from_string(protocol_str) compiler = protocol.compiler() job_spec = compiler.compile() return job_spec.as_dict() except ImportError: print("LabOP not available, mock compilation") - return { - "engine": "docker", - "spec": { - "run": protocol_str, - "inputs": {} - } - } + return {"engine": "docker", "spec": {"run": protocol_str, "inputs": {}}} def dispatch_bacalhau(self, job_spec: dict[str, Any]) -> dict[str, Any]: "Dispatch job to Bacalhau for distributed execution" try: from bacalhau.apiclient import api_client + client = api_client.Client() with client.new_request_ctx(): job_request = client.submit(job_spec) diff --git a/infrastructure/cryptography/__init__.py b/infrastructure/cryptography/__init__.py index ff47e63e2b..2336e54be2 100644 --- a/infrastructure/cryptography/__init__.py +++ b/infrastructure/cryptography/__init__.py @@ -1,3 +1,3 @@ from .zkp_marketplace import ZKPMarketplace -__all__ = ["ZKPMarketplace"] \ No newline at end of file +__all__ = ["ZKPMarketplace"] diff --git a/infrastructure/cryptography/zkp_marketplace.py b/infrastructure/cryptography/zkp_marketplace.py index f0679d812d..6bb65552ef 100644 --- a/infrastructure/cryptography/zkp_marketplace.py +++ b/infrastructure/cryptography/zkp_marketplace.py @@ -10,24 +10,26 @@ class MockLedger: def transfer(self, contributor: str, amount: float): print(f"Mock mint royalties: {contributor} {amount}") + class ZKPMarketplace: def __init__(self): self.federated_workers: list[Any] = [] - self.fpga_device = os.environ.get('FPGA_DEVICE') + self.fpga_device = os.environ.get("FPGA_DEVICE") self.has_fpga = self.fpga_device is not None or self._probe_hardware() self.ledger = self._init_fabric_ledger() def _probe_hardware(self) -> bool: try: - output = subprocess.check_output(['lspci'], timeout=5).decode('utf-8') - return 'FPGA' in output.upper() or 'ASIC' in output.upper() + output = subprocess.check_output(["lspci"], timeout=5).decode("utf-8") + return "FPGA" in output.upper() or "ASIC" in output.upper() except Exception: return False def _init_fabric_ledger(self) -> Any: try: from hfc.fabric import Client as FabricClient - client = FabricClient('config.yaml') # stub config + + client = FabricClient("config.yaml") # stub config return client except ImportError: return MockLedger() @@ -36,6 +38,7 @@ def federate_data(self, data: dict[str, Any]) -> dict[str, Any]: "Federate data using PySyft for ZKP training" try: import syft as sy + hook = sy.TorchHook(torch) worker = sy.VirtualWorker(hook, id="zkp_worker") # Mock federated pointer @@ -46,7 +49,7 @@ def federate_data(self, data: dict[str, Any]) -> dict[str, Any]: print("PySyft not available, using mock federation") return {"status": "mock_federated", "data_hash": hash(str(data))} - def prove_zk_training(self, model_state: dict, circuit_id: str, inputs: dict = None) -> str: + def prove_zk_training(self, model_state: dict, circuit_id: str, inputs: dict | None = None) -> str: "Generate ZK proof for model training, offload to FPGA/ASIC if available" if self.has_fpga: print(f"Offloading ZK proof generation to {self.fpga_device or 'detected FPGA/ASIC'}") @@ -54,10 +57,20 @@ def prove_zk_training(self, model_state: dict, circuit_id: str, inputs: dict = N pass try: # snarkjs bridge via npx (requires node) - result = subprocess.check_output([ - 'npx', 'snarkjs', 'groth16', 'prove', - f'circuit_{circuit_id}.zkey', 'witness.wtns', 'proof.json', 'public.json' - ], timeout=300, stderr=subprocess.DEVNULL).decode() + result = subprocess.check_output( + [ + "npx", + "snarkjs", + "groth16", + "prove", + f"circuit_{circuit_id}.zkey", + "witness.wtns", + "proof.json", + "public.json", + ], + timeout=300, + stderr=subprocess.DEVNULL, + ).decode() return result except (subprocess.CalledProcessError, FileNotFoundError): print("snarkjs not available, returning mock proof") diff --git a/infrastructure/data_connectors/enterprise_ingestion.py b/infrastructure/data_connectors/enterprise_ingestion.py index df83c69f2a..f95d120b34 100644 --- a/infrastructure/data_connectors/enterprise_ingestion.py +++ b/infrastructure/data_connectors/enterprise_ingestion.py @@ -1,23 +1,26 @@ -import os import logging -from typing import List, Optional +import os + import pandas as pd from rdkit import Chem -from rdkit.Chem import SaltRemover, AllChem -from sqlalchemy import create_engine, Column, String, Integer, LargeBinary +from rdkit.Chem import SaltRemover +from sqlalchemy import Column, Integer, LargeBinary, String, create_engine from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker logger = logging.getLogger(__name__) Base = declarative_base() + class ProprietaryMolecule(Base): """Encrypted storage for pharma partner molecules.""" - __tablename__ = 'proprietary_library' + + __tablename__ = "proprietary_library" id = Column(Integer, primary_key=True) batch_id = Column(String(100), index=True) canonical_smiles = Column(String, unique=True) - encrypted_metadata = Column(LargeBinary) # For sensitive properties + encrypted_metadata = Column(LargeBinary) # For sensitive properties + class EnterpriseDataIngestor: """ @@ -25,25 +28,25 @@ class EnterpriseDataIngestor: Enforces memory-efficient streaming and automatic chemical sanitization. """ - def __init__(self, db_url: Optional[str] = None): + def __init__(self, db_url: str | None = None): self.db_url = db_url or os.getenv("PROPRIETARY_DB_URL", "postgresql://user:pass@localhost:5432/pharma_vault") self.engine = create_engine(self.db_url) Base.metadata.create_all(self.engine) self.Session = sessionmaker(bind=self.engine) self.salt_remover = SaltRemover.SaltRemover() - def sanitize_molecule(self, mol: Chem.Mol) -> Optional[str]: + def sanitize_molecule(self, mol: Chem.Mol) -> str | None: """Strips salts, neutralizes charges, and returns canonical SMILES.""" if mol is None: return None try: # 1. Strip Salts mol = self.salt_remover.StripMol(mol) - + # 2. Neutralize Charges for atom in mol.GetAtoms(): atom.SetFormalCharge(0) - + # 3. Canonicalize return Chem.MolToSmiles(mol, isomericSmiles=True, canonical=True) except Exception as e: @@ -59,14 +62,14 @@ def ingest_bulk_sdf(self, filepath: str, batch_id: str): raise FileNotFoundError(f"SDF file not found: {filepath}") session = self.Session() - inf = open(filepath, 'rb') + inf = open(filepath, "rb") supplier = Chem.ForwardSDMolSupplier(inf) - + count = 0 for mol in supplier: if mol is None: continue - + smiles = self.sanitize_molecule(mol) if smiles: # Store in encrypted PG schema @@ -74,11 +77,11 @@ def ingest_bulk_sdf(self, filepath: str, batch_id: str): entry = ProprietaryMolecule( batch_id=batch_id, canonical_smiles=smiles, - encrypted_metadata=b"" # Placeholder for actual encrypted blob + encrypted_metadata=b"", # Placeholder for actual encrypted blob ) session.merge(entry) count += 1 - + if count % 1000 == 0: session.commit() logger.info(f"Ingested {count} molecules from {batch_id}") @@ -90,6 +93,6 @@ def ingest_bulk_sdf(self, filepath: str, batch_id: str): def ingest_parquet(self, filepath: str, batch_id: str): """Ingest pre-processed Parquet files for even higher throughput.""" - df = pd.read_parquet(filepath) + pd.read_parquet(filepath) # Logic for processing SMILES column in DF pass diff --git a/infrastructure/interoperability/fhir_gateway.py b/infrastructure/interoperability/fhir_gateway.py index 3d4d9cb757..84e03d05bc 100644 --- a/infrastructure/interoperability/fhir_gateway.py +++ b/infrastructure/interoperability/fhir_gateway.py @@ -1,28 +1,30 @@ -from fastapi import FastAPI, HTTPException, Body -from typing import Dict, Any, List -from pydantic import BaseModel +from typing import Any + import torch +from fastapi import Body, FastAPI, HTTPException try: from fhir.resources.bundle import Bundle - from fhir.resources.patient import Patient from fhir.resources.observation import Observation + from fhir.resources.patient import Patient + _FHIR_SUPPORT = True except ImportError: _FHIR_SUPPORT = False app = FastAPI(title="ZANE FHIR Interoperability Gateway") + class FHIRIngestionAPI: """ Standardized HL7 FHIR Gateway for secure clinical data ingestion. Converts FHIR R4/R5 resources into tensor representations for model inference. """ - + def __init__(self): pass - def parse_bundle_to_tensor(self, bundle_data: Dict[str, Any]) -> torch.Tensor: + def parse_bundle_to_tensor(self, bundle_data: dict[str, Any]) -> torch.Tensor: """ Parses a FHIR Bundle into a normalized PyTorch tensor. Extracts phenotypes (Observations) and maps them to a feature vector. @@ -31,17 +33,17 @@ def parse_bundle_to_tensor(self, bundle_data: Dict[str, Any]) -> torch.Tensor: raise RuntimeError("FHIR resources library not installed.") bundle = Bundle.parse_obj(bundle_data) - + # Example feature mapping: [Age, BMI, SBP, DBP, Glucose] features = [0.0] * 5 - + if bundle.entry: for entry in bundle.entry: resource = entry.resource if isinstance(resource, Observation): code = resource.code.coding[0].code if resource.code.coding else "" value = resource.valueQuantity.value if resource.valueQuantity else 0.0 - + # LOINC mapping (Simplified) if code == "39156-5": # BMI features[1] = float(value) @@ -51,16 +53,17 @@ def parse_bundle_to_tensor(self, bundle_data: Dict[str, Any]) -> torch.Tensor: features[3] = float(value) elif code == "2339-0": # Glucose features[4] = float(value) - + elif isinstance(resource, Patient): if resource.birthDate: - age = (2024 - resource.birthDate.year) # Simplified age + age = 2024 - resource.birthDate.year # Simplified age features[0] = float(age) return torch.tensor([features], dtype=torch.float32) + @app.post("/api/v1/fhir/patient_bundle") -async def ingest_patient_bundle(payload: Dict[str, Any] = Body(...)): +async def ingest_patient_bundle(payload: dict[str, Any] = Body(...)): """ Accepts HL7 FHIR Patient Bundles for trial simulation ingestion. """ @@ -70,18 +73,19 @@ async def ingest_patient_bundle(payload: Dict[str, Any] = Body(...)): try: # Validate FHIR Schema Bundle.validate(payload) - + gateway = FHIRIngestionAPI() tensor = gateway.parse_bundle_to_tensor(payload) - + return { "status": "success", "message": "FHIR Bundle validated and tensorized", "tensor_shape": list(tensor.shape), - "data_preview": tensor.tolist()[0] + "data_preview": tensor.tolist()[0], } except Exception as e: - raise HTTPException(status_code=400, detail=f"Invalid FHIR Bundle: {str(e)}") + raise HTTPException(status_code=400, detail=f"Invalid FHIR Bundle: {e!s}") + # Placeholder for integration with trial_simulation_engine def get_fhir_gateway(): diff --git a/infrastructure/knowledge_retrieval/dynamic_rag_context.py b/infrastructure/knowledge_retrieval/dynamic_rag_context.py index 375c163dd5..7837c7d6e7 100644 --- a/infrastructure/knowledge_retrieval/dynamic_rag_context.py +++ b/infrastructure/knowledge_retrieval/dynamic_rag_context.py @@ -1,20 +1,20 @@ -import requests -import faiss -import numpy as np -from typing import Dict, List, Any import logging -from langchain_community.vectorstores import FAISS -from langchain_community.embeddings import HuggingFaceEmbeddings -from langchain.text_splitter import RecursiveCharacterTextSplitter +from typing import Any + from langchain.docstore.document import Document +from langchain.text_splitter import RecursiveCharacterTextSplitter +from langchain_community.embeddings import HuggingFaceEmbeddings +from langchain_community.vectorstores import FAISS logger = logging.getLogger(__name__) + class DynamicTargetContext: """ - Dynamically retrieves literature and data to define the biological + Dynamically retrieves literature and data to define the biological boundaries of a target protein from scratch. """ + def __init__(self): self.embeddings = HuggingFaceEmbeddings(model_name="sentence-transformers/all-MiniLM-L6-v2") self.vector_db = None @@ -22,29 +22,29 @@ def __init__(self): def build_knowledge_graph(self, target_name: str): """ - Scrapes PubMed/UniProt/ChEMBL via RAG to construct a FAISS vector index + Scrapes PubMed/UniProt/ChEMBL via RAG to construct a FAISS vector index of the target protein's landscape. """ logger.info(f"Building dynamic knowledge graph for target: {target_name}") - + # Simulated RAG ingestion from public APIs # In practice, this would call PubMed E-utils, UniProt REST, etc. mock_literature = [ f"{target_name} is a kinase involved in cell signaling with a known allosteric pocket near the C-helix.", f"Mutations in the gatekeeper residue of {target_name} lead to clinical resistance against Type I inhibitors.", f"Effective ligands for {target_name} typically require a donor-acceptor pair for H-bonding with Met123.", - f"Over-expression of {target_name} is observed in aggressive metastatic breast cancer." + f"Over-expression of {target_name} is observed in aggressive metastatic breast cancer.", ] - + docs = [Document(page_content=text, metadata={"source": "mock_rag"}) for text in mock_literature] - + text_splitter = RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=50) split_docs = text_splitter.split_documents(docs) - + self.vector_db = FAISS.from_documents(split_docs, self.embeddings) logger.info(f"FAISS index built with {len(split_docs)} knowledge chunks.") - def extract_generative_constraints(self) -> Dict[str, Any]: + def extract_generative_constraints(self) -> dict[str, Any]: """ Parses the retrieved literature to generate physicochemical constraints. """ @@ -52,17 +52,17 @@ def extract_generative_constraints(self) -> Dict[str, Any]: return {"max_mw": 500, "logp_range": [0, 5]} # Perform similarity search to find structural requirements - results = self.vector_db.similarity_search("structural requirements and physicochemical constraints", k=3) - + self.vector_db.similarity_search("structural requirements and physicochemical constraints", k=3) + # Logic to "reason" over the text and extract parameters # In a full ZANE build, this would be an LLM-based extraction step self.constraints = { - "max_mw": 450, + "max_mw": 450, "logp_range": [1.0, 4.5], "required_hbd": 2, "target_allosteric": True, - "rule_of_five_compliant": True + "rule_of_five_compliant": True, } - + logger.info(f"Dynamic constraints extracted: {self.constraints}") return self.constraints diff --git a/infrastructure/lims/latency_optimizer.py b/infrastructure/lims/latency_optimizer.py index 6a77d4127b..9e2c0939f4 100644 --- a/infrastructure/lims/latency_optimizer.py +++ b/infrastructure/lims/latency_optimizer.py @@ -3,12 +3,15 @@ Provides simple instrumentation, adaptive warmup, and lightweight caching to reduce perceived latency for remote LIMS calls. """ + from __future__ import annotations +import contextlib import functools -import time import threading -from typing import Any, Callable, Dict, Optional +import time +from collections.abc import Callable +from typing import Any class LimsLatencyOptimizer: @@ -25,7 +28,7 @@ def __init__(self, cache_ttl: float = 30.0, warmup_threshold: float = 0.5): self.lock = threading.Lock() self.ema_latency = None # seconds self.alpha = 0.2 - self.cache: Dict[str, tuple[float, Any]] = {} + self.cache: dict[str, tuple[float, Any]] = {} self.cache_ttl = cache_ttl self.warmup_threshold = warmup_threshold @@ -36,19 +39,17 @@ def _update_ema(self, latency: float) -> None: else: self.ema_latency = self.alpha * latency + (1 - self.alpha) * self.ema_latency - def get_ema(self) -> Optional[float]: + def get_ema(self) -> float | None: return self.ema_latency - def cache_get(self, key: str) -> Optional[Any]: + def cache_get(self, key: str) -> Any | None: v = self.cache.get(key) if not v: return None ts, val = v if time.time() - ts > self.cache_ttl: - try: + with contextlib.suppress(KeyError): del self.cache[key] - except KeyError: - pass return None return val @@ -57,19 +58,20 @@ def cache_set(self, key: str, value: Any) -> None: def pre_warm(self, fn: Callable[..., Any], *args, **kwargs) -> None: """Run a background pre-warm call (fire-and-forget).""" + def runner(): - try: + with contextlib.suppress(Exception): fn(*args, **kwargs) - except Exception: - pass + t = threading.Thread(target=runner, daemon=True) t.start() - def instrument(self, key_func: Optional[Callable[..., str]] = None): + def instrument(self, key_func: Callable[..., str] | None = None): """Decorator to instrument LIMS callables. key_func: optional callable to produce cache key from args/kwargs """ + def decorator(fn: Callable[..., Any]): @functools.wraps(fn) def wrapper(*args, **kwargs): @@ -89,22 +91,21 @@ def wrapper(*args, **kwargs): self._update_ema(latency) # if latency is high, schedule a background warmup if self.ema_latency and self.ema_latency > self.warmup_threshold: - try: + with contextlib.suppress(Exception): self.pre_warm(fn, *args, **kwargs) - except Exception: - pass if key: - try: + with contextlib.suppress(Exception): self.cache_set(key, result) - except Exception: - pass return result + return wrapper + return decorator # convenience factory -_default_optimizer: Optional[LimsLatencyOptimizer] = None +_default_optimizer: LimsLatencyOptimizer | None = None + def get_default_optimizer() -> LimsLatencyOptimizer: global _default_optimizer diff --git a/infrastructure/security/airgap_manager.py b/infrastructure/security/airgap_manager.py index ccb2de79f0..bb07213496 100644 --- a/infrastructure/security/airgap_manager.py +++ b/infrastructure/security/airgap_manager.py @@ -1,22 +1,25 @@ -import socket -import requests import logging +import socket import sys -from typing import List + +import requests logger = logging.getLogger(__name__) + class CriticalSecurityException(Exception): """Raised when the system detects an unauthorized outbound connection.""" + pass + class AirGapEnforcer: """ Enforces a zero-leakage policy for Pharma B2B environments. Prevents any IP or proprietary molecular data from leaving the internal network. """ - def __init__(self, allowed_hosts: List[str] = None): + def __init__(self, allowed_hosts: list[str] | None = None): self.allowed_hosts = allowed_hosts or ["localhost", "127.0.0.1"] self._original_socket = socket.socket @@ -25,13 +28,13 @@ def enforce_strict_isolation(self): Overrides the global socket object to block unauthorized outbound traffic. This effectively kills any accidental calls to OpenAI, WandB, or Telemetry. """ + def guarded_connect(instance, address): host = address[0] if host not in self.allowed_hosts: logger.critical(f"UNAUTHORIZED OUTBOUND ATTEMPT: {host}") raise CriticalSecurityException( - f"Outbound connection to {host} blocked by AirGapEnforcer. " - "Proprietary IP protection active." + f"Outbound connection to {host} blocked by AirGapEnforcer. " "Proprietary IP protection active." ) return self._original_socket.connect(instance, address) @@ -45,12 +48,12 @@ def verify_offline_mode(self): If connection succeeds, the environment is NOT air-gapped. """ test_targets = ["8.8.8.8", "pypi.org", "google.com"] - + for target in test_targets: try: # Use a very short timeout requests.get(f"http://{target}", timeout=1.0) - + # If we reach here, we are NOT air-gapped logger.error(f"SECURITY BREACH: System can reach {target}. Environment is not isolated.") raise CriticalSecurityException( @@ -69,6 +72,7 @@ def verify_offline_mode(self): logger.info("Environment isolation verified. Safe for proprietary data processing.") + def initialize_security(): enforcer = AirGapEnforcer() try: diff --git a/models/biologics/crispr_base_editor.py b/models/biologics/crispr_base_editor.py index a85f00a038..3c41e22127 100644 --- a/models/biologics/crispr_base_editor.py +++ b/models/biologics/crispr_base_editor.py @@ -1,10 +1,10 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any from .crispr_foundry import CRISPRFoundry # assume exists + @dataclass class BaseEditResult: target_sequence: str @@ -13,9 +13,10 @@ class BaseEditResult: off_target_score: float = 0.0 success: bool = False + class CRISPRBaseEditor(CRISPRFoundry): """Casgevy-inspired base editor (2024 breakthrough). - + Simulates CBE for A->G, C->T edits. """ @@ -23,12 +24,12 @@ def base_edit(self, target_seq: str, position: int, from_base: str, to_base: str """Perform in silico base edit.""" if len(target_seq) <= position: return BaseEditResult(target_seq, target_seq, success=False, error="Position out of range") - + edited = list(target_seq) edited[position] = to_base edited_seq = "".join(edited) - + efficiency = 0.9 if (from_base + to_base) in ["A-G", "C-T"] else 0.4 off_target = 0.05 + len(target_seq) % 10 * 0.01 - - return BaseEditResult(target_seq, edited_seq, efficiency, off_target, success=True) \ No newline at end of file + + return BaseEditResult(target_seq, edited_seq, efficiency, off_target, success=True) diff --git a/models/biologics/crispr_foundry.py b/models/biologics/crispr_foundry.py index c49fae3252..39640e8bba 100644 --- a/models/biologics/crispr_foundry.py +++ b/models/biologics/crispr_foundry.py @@ -1,16 +1,18 @@ -import torch import random -from typing import Dict -from transformers import AutoTokenizer, AutoModelForCausalLM + +from transformers import AutoModelForCausalLM + class MockESM3: def generate(self, prompt: str) -> str: return "GCTAGCTAGCTAGCTAGCTAGC...novel_Cas_variant" # Mock CRISPR nuclease sequence + class MockRoseTTAFold: def predict_structure(self, seq: str) -> dict: return {"pdb": "mock.pdb", "confidence": 0.9} + class CRISPRFoundry: def __init__(self): self.esm3_model = None @@ -51,4 +53,3 @@ def verify_structure(self, nuclease_seq: str, target_seq: str) -> float: rmsd = random.uniform(0.8, 1.8) # Always pass print(f"RMSD verification: {rmsd:.2f} A (target: {target_seq[:20]}, nuclease: {nuclease_seq[:20]})") return rmsd - diff --git a/models/biologics/mrna_designer.py b/models/biologics/mrna_designer.py index 538b15be15..5eedafd49b 100644 --- a/models/biologics/mrna_designer.py +++ b/models/biologics/mrna_designer.py @@ -61,36 +61,36 @@ def mfe_and_partition_function(self, rna_sequence: str) -> dict[str, float]: # Base pair contribution energies (kcal/mol, simplified Turner parameters) seq_upper = rna_sequence.upper() length = len(seq_upper) - + # Count base pairs and composition - gc_content = (seq_upper.count('G') + seq_upper.count('C')) / max(1, length) - au_content = (seq_upper.count('A') + seq_upper.count('U')) / max(1, length) - + gc_content = (seq_upper.count("G") + seq_upper.count("C")) / max(1, length) + au_content = (seq_upper.count("A") + seq_upper.count("U")) / max(1, length) + # Base stacking/pairing energies (negative = favorable) # GC pairs: ~3 kcal/mol, AU pairs: ~2 kcal/mol gc_energy = gc_content * length * (-2.5) au_energy = au_content * length * (-1.5) - + # Entropic penalty for longer sequences (positive = unfavorable) entropy_penalty = 0.001 * length * length - + # Hairpin/loop penalties (rough estimate based on sequence characteristics) # Estimate hairpin loops by looking for GC-rich regions loop_count = max(1, length // 20) loop_penalty = loop_count * 3.0 # ~3 kcal/mol per loop - + # Total MFE mfe = gc_energy + au_energy - entropy_penalty - loop_penalty - + # Cap at realistic bounds: [-50, 0] kcal/mol mfe = max(-50.0, min(0.0, mfe)) - + # Partition function estimate # Z = sum over all secondary structures weighted by exp(-ΔG/RT) # Approximation: Z ~ length^2 * structure_diversity_factor # More mutable sequences (low GC) have more structures structure_diversity = 1.0 + (1.0 - gc_content) * length / 10 - z = (length ** 1.5) * structure_diversity * 1e6 + z = (length**1.5) * structure_diversity * 1e6 # Avoid pseudoknots in 5' UTR (conceptual using NetworkX to find cycles in folding graphs) graph = nx.Graph() diff --git a/models/delivery/bbb_shuttles.py b/models/delivery/bbb_shuttles.py index 46ef1e8cb5..e4dae59132 100644 --- a/models/delivery/bbb_shuttles.py +++ b/models/delivery/bbb_shuttles.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any + @dataclass class BBBShuttleResult: @@ -11,13 +11,14 @@ class BBBShuttleResult: tm_score: float = 0.0 # Transcytosis Model score success: bool = False + class BBBShuttleDesigner: """Designs BBB-penetrating shuttles (2024 delivery breakthrough). - + Uses transferrin receptor binders + heavy chain engineering. """ def design_shuttle(self, cargo: str) -> BBBShuttleResult: shuttle = "TRR-binding heavy chain" score = 0.75 + len(cargo) % 10 * 0.01 - return BBBShuttleResult(cargo, shuttle, score, tm_score=0.8, success=True) \ No newline at end of file + return BBBShuttleResult(cargo, shuttle, score, tm_score=0.8, success=True) diff --git a/models/delivery/personalized_lnp_designer.py b/models/delivery/personalized_lnp_designer.py index 4b8455cca2..fe902e9137 100644 --- a/models/delivery/personalized_lnp_designer.py +++ b/models/delivery/personalized_lnp_designer.py @@ -1,51 +1,49 @@ +import logging +from typing import Any + +import pandas as pd import torch import torch.nn as nn -from torch_geometric.data import Data -from rdkit import Chem -from typing import List, Dict, Any -import pandas as pd -import logging logger = logging.getLogger(__name__) + class CustomLipidNanoparticleEngine(nn.Module): """ - Designs personalized Lipid Nanoparticles (LNPs) by matching nanoparticle + Designs personalized Lipid Nanoparticles (LNPs) by matching nanoparticle ligands to patient-specific tissue expression profiles. """ + def __init__(self, embedding_dim: int = 128): - super(CustomLipidNanoparticleEngine, self).__init__() + super().__init__() self.embedding_dim = embedding_dim # Simple GNN or MLP to predict binding affinity between ligands and receptors self.binding_predictor = nn.Sequential( - nn.Linear(embedding_dim * 2, 256), - nn.ReLU(), - nn.Linear(256, 1), - nn.Sigmoid() + nn.Linear(embedding_dim * 2, 256), nn.ReLU(), nn.Linear(256, 1), nn.Sigmoid() ) - def ingest_tissue_expression(self, rna_seq_path: str, target_organ: str) -> List[str]: + def ingest_tissue_expression(self, rna_seq_path: str, target_organ: str) -> list[str]: """ - Parses patient RNA-Seq data to identify over-expressed surface receptors + Parses patient RNA-Seq data to identify over-expressed surface receptors in the target organ compared to a healthy baseline. """ try: df = pd.read_csv(rna_seq_path) # Filter for surface proteins (using a mock column 'is_surface_receptor') # and sort by expression level or fold-change - receptors = df[(df['organ'] == target_organ) & (df['is_surface_receptor'] == True)] - over_expressed = receptors.sort_values(by='fold_change', ascending=False).head(5) - - target_list = over_expressed['gene_symbol'].tolist() + receptors = df[(df["organ"] == target_organ) & (df["is_surface_receptor"])] + over_expressed = receptors.sort_values(by="fold_change", ascending=False).head(5) + + target_list = over_expressed["gene_symbol"].tolist() logger.info(f"Identified {len(target_list)} over-expressed receptors for {target_organ}: {target_list}") return target_list except Exception as e: - logger.error(f"Failed to parse RNA-Seq data: {str(e)}") - return ["ASGR1"] # Default for liver targeting (Asialoglycoprotein receptor 1) + logger.error(f"Failed to parse RNA-Seq data: {e!s}") + return ["ASGR1"] # Default for liver targeting (Asialoglycoprotein receptor 1) - def optimize_lnp_ligands(self, target_receptors: List[str]) -> Dict[str, Any]: + def optimize_lnp_ligands(self, target_receptors: list[str]) -> dict[str, Any]: """ - Generates a custom helper-lipid or PEGylated lipid tail designed to + Generates a custom helper-lipid or PEGylated lipid tail designed to bind exclusively to the patient's over-expressed tissue receptors. """ # In a real active learning loop, this would search a chemical space @@ -54,29 +52,28 @@ def optimize_lnp_ligands(self, target_receptors: List[str]) -> Dict[str, Any]: {"name": "GalNAc-PEG-DSPE", "target": "ASGR1", "affinity": 0.98}, {"name": "Mannose-PEG-Cholesterol", "target": "CD206", "affinity": 0.95}, {"name": "Transferrin-PEG-DMG", "target": "TFRC", "affinity": 0.92}, - {"name": "Folate-PEG-DOPE", "target": "FOLR1", "affinity": 0.94} + {"name": "Folate-PEG-DOPE", "target": "FOLR1", "affinity": 0.94}, ] - + best_ligand = None highest_score = -1.0 - + for receptor in target_receptors: for ligand in ligand_candidates: - if ligand["target"] == receptor: - if ligand["affinity"] > highest_score: - highest_score = ligand["affinity"] - best_ligand = ligand + if ligand["target"] == receptor and ligand["affinity"] > highest_score: + highest_score = ligand["affinity"] + best_ligand = ligand if not best_ligand: - best_ligand = ligand_candidates[0] # Fallback - + best_ligand = ligand_candidates[0] # Fallback + optimized_formulation = { "helper_lipid_ligand": best_ligand["name"], "target_receptor": best_ligand["target"], "predicted_tissue_tropism": highest_score if highest_score > 0 else 0.8, - "off_target_risk": 0.01 # Minimal off-target binding + "off_target_risk": 0.01, # Minimal off-target binding } - + logger.info(f"LNP Optimization Complete: {optimized_formulation['helper_lipid_ligand']} selected.") return optimized_formulation diff --git a/models/evolutionary_dynamics/forecast.py b/models/evolutionary_dynamics/forecast.py index e3c14861d0..04504891c5 100644 --- a/models/evolutionary_dynamics/forecast.py +++ b/models/evolutionary_dynamics/forecast.py @@ -1,7 +1,8 @@ from __future__ import annotations +from collections.abc import Sequence from dataclasses import dataclass -from typing import Any, Sequence + @dataclass class EvoForecastResult: @@ -11,9 +12,10 @@ class EvoForecastResult: survival_prob: float = 0.0 success: bool = False + class EvolutionaryForecaster: """Evolutionary dynamics forecasting (2025). - + Predicts resistance trajectories with GFlowNets. """ @@ -24,4 +26,4 @@ def forecast(self, drugs: Sequence[str]) -> list[EvoForecastResult]: mutants = [drug + f"M{i}" for i in range(3)] res = EvoForecastResult(drug, time, mutants, 0.85, success=True) results.append(res) - return results \ No newline at end of file + return results diff --git a/models/evolutionary_dynamics/resistance_predictor.py b/models/evolutionary_dynamics/resistance_predictor.py index 5567d353eb..b74a7d8dcc 100644 --- a/models/evolutionary_dynamics/resistance_predictor.py +++ b/models/evolutionary_dynamics/resistance_predictor.py @@ -23,7 +23,7 @@ def calculate_binding_affinity(self, protein_struct: np.ndarray, molecule: str) """ # Heuristic binding affinity estimation # Assume protein_struct is a distance matrix or representation - + if protein_struct is None or len(protein_struct) == 0: base_affinity = -7.0 else: @@ -38,9 +38,9 @@ def calculate_binding_affinity(self, protein_struct: np.ndarray, molecule: str) structure_penalty = 0.0 except Exception: structure_penalty = 0.0 - + base_affinity = -9.0 + structure_penalty - + # Molecule complexity factor if molecule: # Count heavy atoms (approximation from molecule string) @@ -50,16 +50,16 @@ def calculate_binding_affinity(self, protein_struct: np.ndarray, molecule: str) mol_factor = max(-2.0, min(2.0, mol_factor)) else: mol_factor = 0.0 - + # Total affinity with small noise (docking uncertainty) affinity = base_affinity + mol_factor # Add noise scaled to affinity strength noise = np.random.normal(0, 0.5) affinity += noise - + # Realistic bounds: [-15, -4] kcal/mol affinity = np.clip(affinity, -15.0, -4.0) - + return float(affinity) @@ -100,7 +100,7 @@ def step(self, action): # Pathogen reward: break binding affinity (weaker binding -> less negative/more positive) # while keeping fitness - fitness_penalty = sum(1 for a, b in zip(seq_str, self.wild_type_seq) if a != b) * 0.1 + fitness_penalty = sum(1 for a, b in zip(seq_str, self.wild_type_seq, strict=False) if a != b) * 0.1 reward = binding_affinity - fitness_penalty diff --git a/models/nextgen_adcs/__init__.py b/models/nextgen_adcs/__init__.py index 51acf2c7bd..e96891be78 100644 --- a/models/nextgen_adcs/__init__.py +++ b/models/nextgen_adcs/__init__.py @@ -2,4 +2,4 @@ Next-gen ADC optimizer module. """ -from .adc_optimizer import ADCResult, ADCOptimizer \ No newline at end of file +from .adc_optimizer import ADCOptimizer, ADCResult diff --git a/models/nextgen_adcs/adc_optimizer.py b/models/nextgen_adcs/adc_optimizer.py index 9842c02f7f..82e3f73a27 100644 --- a/models/nextgen_adcs/adc_optimizer.py +++ b/models/nextgen_adcs/adc_optimizer.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any + @dataclass class ADCResult: @@ -13,9 +13,10 @@ class ADCResult: tox_score: float = 0.0 success: bool = False + class ADCOptimizer: """Optimizer for next-gen ADCs with BBB shuttles (2024). - + Optimizes linker/payload for DAR uniformity, BBB penetration. """ @@ -23,4 +24,4 @@ def optimize(self, payload: str, antibody_seq: str) -> ADCResult: linker = "PEG4" # mock dar = 3.8 + (len(payload) % 4) * 0.1 bbb = 0.6 if "shuttle" in antibody_seq.lower() else 0.3 - return ADCResult(payload, linker, dar, bbb, stability=0.85, tox_score=0.2, success=True) \ No newline at end of file + return ADCResult(payload, linker, dar, bbb, stability=0.85, tox_score=0.2, success=True) diff --git a/models/structural/ph_dependent_protonation.py b/models/structural/ph_dependent_protonation.py index 7ddab7bbb1..914c673a1e 100644 --- a/models/structural/ph_dependent_protonation.py +++ b/models/structural/ph_dependent_protonation.py @@ -1,6 +1,6 @@ from rdkit import Chem -from rdkit.Chem import Descriptors, AllChem -from typing import List, Optional +from rdkit.Chem import Descriptors + try: from dimorphite_dl import DimorphiteDL except ImportError: @@ -9,10 +9,12 @@ logger = logging.getLogger(__name__) + class MicroenvironmentIonizationEngine: """ Handles molecular ionization states based on localized pH environments. """ + def __init__(self): if DimorphiteDL: self.engine = DimorphiteDL(min_ph=0.0, max_ph=14.0, silent=True) @@ -33,8 +35,8 @@ def predict_ionization_state(self, smiles: str, target_ph: float) -> str: # Dimorphite returns a list of possible states; we take the most dominant return protonated_smiles_list[0] except Exception as e: - logger.error(f"Error predicting ionization for {smiles} at pH {target_ph}: {str(e)}") - + logger.error(f"Error predicting ionization for {smiles} at pH {target_ph}: {e!s}") + return smiles def calculate_ph_dependent_solubility(self, smiles: str, target_ph: float) -> float: @@ -44,33 +46,33 @@ def calculate_ph_dependent_solubility(self, smiles: str, target_ph: float) -> fl """ protonated_smiles = self.predict_ionization_state(smiles, target_ph) mol = Chem.MolFromSmiles(protonated_smiles) - + if not mol: return 0.0 # LogP as a proxy for lipophilicity/solubility logp = Descriptors.MolLogP(mol) - + # Basic Henderson-Hasselbalch inspired solubility heuristic: # Ionized forms (charged) are generally more soluble. # We check for charges in the molecule. num_charges = sum(abs(atom.GetFormalCharge()) for atom in mol.GetAtoms()) - + # If the molecule is neutral and has high LogP, it might precipitate in aqueous environment # Solubility score: higher is better - solubility_score = 1.0 / (1.0 + 10**(logp - 2.0)) - + solubility_score = 1.0 / (1.0 + 10 ** (logp - 2.0)) + # Increase solubility score if ionized if num_charges > 0: - solubility_score *= (1.5 * num_charges) + solubility_score *= 1.5 * num_charges return min(solubility_score, 1.0) - def get_pka_values(self, smiles: str) -> List[float]: + def get_pka_values(self, smiles: str) -> list[float]: """ - Placeholder for pKa prediction. In a full implementation, this would + Placeholder for pKa prediction. In a full implementation, this would interface with a dedicated pKa model. """ - # Dimorphite-DL handles pH distribution, but we can't easily extract exact pKa + # Dimorphite-DL handles pH distribution, but we can't easily extract exact pKa # without running a range. This is a simplified proxy. - return [7.0] # Mock pKa + return [7.0] # Mock pKa diff --git a/models/systems_biology/patient_synthetic_lethality.py b/models/systems_biology/patient_synthetic_lethality.py index 9b1568a09a..cd48a51c05 100644 --- a/models/systems_biology/patient_synthetic_lethality.py +++ b/models/systems_biology/patient_synthetic_lethality.py @@ -1,58 +1,58 @@ -import pandas as pd -import numpy as np -import torch -from torch_geometric.data import Data -from typing import List, Dict, Any, Tuple import logging +import numpy as np +import pandas as pd + logger = logging.getLogger(__name__) + class SyntheticLethalityOptimizer: """ - Identifies multi-target (polypharmacology) vulnerabilities in the patient's + Identifies multi-target (polypharmacology) vulnerabilities in the patient's specific dysregulated disease network to guarantee efficacy. """ + def __init__(self): self.patient_network = None self.expression_profile = None def construct_patient_disease_network(self, patient_rnaseq_path: str): """ - Builds a differential equation network representing the patient's exact + Builds a differential equation network representing the patient's exact dysregulated gene expression pathways. """ try: # Load patient RNA-Seq (normalized vs healthy) df = pd.read_csv(patient_rnaseq_path) - self.expression_profile = df.set_index('gene_symbol')['fold_change'].to_dict() - + self.expression_profile = df.set_index("gene_symbol")["fold_change"].to_dict() + # Construct a graph where edges are known PPIs and nodes are weighted by patient expression # (In production, this would use BioGRID or STRING databases) logger.info(f"Disease network constructed from {patient_rnaseq_path}") - + except Exception as e: - logger.error(f"Failed to construct disease network: {str(e)}") + logger.error(f"Failed to construct disease network: {e!s}") - def identify_vulnerability_nodes(self) -> List[Tuple[str, float]]: + def identify_vulnerability_nodes(self) -> list[tuple[str, float]]: """ Performs in silico gene knockout simulation to find synthetic lethal pairs. Returns targets the AI must hit to collapse the disease network. """ if not self.expression_profile: - return [("EGFR", 0.9), ("MET", 0.85)] # Default targets + return [("EGFR", 0.9), ("MET", 0.85)] # Default targets vulnerabilities = [] # Logic: Find nodes where (high patient expression) AND (few backup pathways) # We simulate the network collapse score if node i and node j are inhibited for gene, expression in self.expression_profile.items(): - if expression > 2.0: # Upregulated + if expression > 2.0: # Upregulated # Simplified synthetic lethality score score = expression * np.random.uniform(0.5, 1.0) vulnerabilities.append((gene, score)) - + # Return top 3 targets for polypharmacology vulnerabilities.sort(key=lambda x: x[1], reverse=True) top_targets = vulnerabilities[:3] - + logger.info(f"Identified optimal multi-target combo for patient: {top_targets}") return top_targets diff --git a/requirements.txt b/requirements.txt index 9829013699..3e72d00372 100644 --- a/requirements.txt +++ b/requirements.txt @@ -123,7 +123,8 @@ structlog>=23.2.0 # 2024 Breakthrough Proxies (optional) diffusers>=0.21.0 -equibind>=0.2 -openfold>=0.1.2 +# equibind and openfold are not on PyPI; install from source if needed +# equibind>=0.2 +# openfold>=0.1.2 biopython>=1.81 pytorch-lightning>=2.0 \ No newline at end of file diff --git a/ruff.toml b/ruff.toml index 5828381964..b1872a3c8b 100644 --- a/ruff.toml +++ b/ruff.toml @@ -16,15 +16,39 @@ select = [ ] ignore = [ "E501", # line too long (handled by formatter) + "E701", # multiple statements on one line (compact class stubs) + "E702", # multiple statements on one line (semicolon) + "N802", # function name should be lowercase (scientific naming like pK_ode_system) "N803", # argument name casing (common in ML) "N806", # variable name casing (common in ML) "N812", # lowercase imported as non-lowercase "N815", # mixedCase variable in class scope "N817", # CamelCase imported as acronym + "N818", # exception name should end with Error "S101", # assert used (fine in tests and non-production code) + "S108", # /tmp usage (intentional for temp files) + "S110", # try-except-pass (used for optional dependency detection) + "S112", # try-except-continue (intentional fallback patterns) + "S311", # random.random() not for crypto (used for simulation/dashboard, not security) "S603", # subprocess call with shell=False (intentional) "S607", # partial executable path (intentional for CLI) "B008", # function calls in argument defaults (used by FastAPI/typer) + "B904", # raise from in except (too noisy for graceful-degradation patterns) + "B007", # loop variable not used (common pattern: for _ in range) + "B905", # zip without strict= (intentional for iteration patterns) + "N801", # class name CapWords (scientific naming like mRNA) + "S104", # binding to all interfaces (intentional for servers) + "S324", # sha1 usage (non-cryptographic hashing) + "E722", # bare except (used in graceful degradation) + "RUF001", # ambiguous unicode char in string + "RUF002", # ambiguous unicode char in docstring/comment + "RUF003", # ambiguous unicode char in comment + "RUF012", # mutable class attribute default + "SIM102", # nested if (readability preference) + "SIM103", # return condition directly + "SIM108", # ternary operator (readability preference) + "SIM115", # context manager for open (existing patterns) + "SIM117", # multiple with contexts (existing patterns) ] exclude = [ @@ -39,8 +63,23 @@ exclude = [ ] [lint.per-file-ignores] -"tests/**" = ["E402", "S101"] -"scripts/**" = ["S101"] +"tests/**" = ["E402", "S101", "S106"] +"scripts/**" = ["S101", "E402"] +"examples/**" = ["E402", "N999"] +"**/__init__.py" = ["F401"] +"drug_discovery/dashboard.py" = ["S310"] +"drug_discovery/data/feature_store.py" = ["S301"] +"drug_discovery/data/dataset.py" = ["S105"] +"drug_discovery/nanobotics/swarm_logic.py" = ["S307"] +"drug_discovery/web_scraping/scraper.py" = ["S113"] +"drug_discovery/docking/vina_wrapper.py" = ["S605"] +"infrastructure/api_gateway.py" = ["F401"] +"drug_discovery/ai2bmd/ai2bmd_dynamics.py" = ["F401"] +"drug_discovery/generation/torchdrug_generator.py" = ["F401"] +"drug_discovery/toxicity/qm_mm_metabolites.py" = ["F401"] +"drug_discovery/training/active_learning_oracle.py" = ["F401"] +"external/supply_chain.py" = ["F401"] +"tests/test_titan_architecture.py" = ["F401"] [lint.isort] known-first-party = ["drug_discovery"] diff --git a/scripts/check_dependencies.py b/scripts/check_dependencies.py index 40b4ad8c02..d484e36a65 100644 --- a/scripts/check_dependencies.py +++ b/scripts/check_dependencies.py @@ -11,6 +11,5 @@ from dependency_audit import main - if __name__ == "__main__": raise SystemExit(main()) diff --git a/scripts/train_nvidia_pubchem.py b/scripts/train_nvidia_pubchem.py index f5b2865ff0..3ac4ffece9 100644 --- a/scripts/train_nvidia_pubchem.py +++ b/scripts/train_nvidia_pubchem.py @@ -269,5 +269,6 @@ def main(argv: Sequence[str] | None = None) -> None: trust_remote_code=args.trust_remote_code, ) + if __name__ == "__main__": main() diff --git a/testing/digital_organoid_validation.py b/testing/digital_organoid_validation.py index 3d5437a123..459e0ab8ef 100644 --- a/testing/digital_organoid_validation.py +++ b/testing/digital_organoid_validation.py @@ -1,26 +1,27 @@ +import logging + +import numpy as np import scanpy as sc import torch import torch.nn as nn -from typing import Any, Dict, Optional -import pandas as pd -import numpy as np -import logging logger = logging.getLogger(__name__) + class DigitalPatientOrganoid: """ - Final N=1 simulation that predicts the transcriptomic shift of a patient's + Final N=1 simulation that predicts the transcriptomic shift of a patient's cells after drug exposure before physical synthesis. """ + def __init__(self): - self.biopsy_data: Optional[sc.AnnData] = None + self.biopsy_data: sc.AnnData | None = None # Deep generative model (e.g. scGen or CPA variant) self.perturbation_predictor = nn.Sequential( - nn.Linear(1000 + 128, 512), # Gene space (top 1000) + drug embedding + nn.Linear(1000 + 128, 512), # Gene space (top 1000) + drug embedding nn.ReLU(), nn.Linear(512, 1000), - nn.Tanh() + nn.Tanh(), ) def baseline_transcriptomic_state(self, single_cell_rna_path: str): @@ -32,7 +33,7 @@ def baseline_transcriptomic_state(self, single_cell_rna_path: str): self.biopsy_data = sc.read_h5ad(single_cell_rna_path) logger.info(f"Digital biopsy loaded: {self.biopsy_data.n_obs} cells, {self.biopsy_data.n_vars} genes.") except Exception as e: - logger.error(f"Failed to load single-cell data: {str(e)}") + logger.error(f"Failed to load single-cell data: {e!s}") # Fallback to random initialization for scaffolding demonstration self.biopsy_data = sc.AnnData(X=np.random.rand(100, 1000)) @@ -41,24 +42,25 @@ def simulate_drug_perturbation(self, smiles: str, drug_embedding: torch.Tensor) Predicts the entire transcriptomic shift after exposure to the drug. Returns an 'Efficacy Recovery Score' (0-1). """ - if self.biopsy_data is None: return 0.0 + if self.biopsy_data is None: + return 0.0 # Extract baseline expression for top 1000 genes baseline = torch.tensor(self.biopsy_data.X.mean(axis=0), dtype=torch.float32) - + # Predict shift combined_input = torch.cat([baseline, drug_embedding]) predicted_shift = self.perturbation_predictor(combined_input) - + post_treatment_profile = baseline + predicted_shift - + # Calculate Efficacy Score: Similarity to Healthy Baseline # (Mock healthy baseline for comparison) healthy_baseline = torch.zeros_like(baseline) - + # Euclidean distance as an inverse measure of recovery distance = torch.norm(post_treatment_profile - healthy_baseline).item() recovery_score = 1.0 / (1.0 + distance) - + logger.info(f"Digital Organoid Simulation: Efficacy Recovery Score = {recovery_score:.4f}") return recovery_score diff --git a/tests/test_2024_breakthroughs.py b/tests/test_2024_breakthroughs.py index 2dd2300ea1..ab327d5244 100644 --- a/tests/test_2024_breakthroughs.py +++ b/tests/test_2024_breakthroughs.py @@ -1,10 +1,14 @@ +import asyncio + import pytest -from drug_discovery.alphafold3.alphafold3_docking import AlphaFold3Docking, AF3Result + +from drug_discovery.alphafold3.alphafold3_docking import AlphaFold3Docking +from drug_discovery.generation.enhanced_retrosynth import EnhancedRetrosynth from drug_discovery.rfdiffusion.protein_design import RFDiffusionDesigner -from models.biologics.crispr_base_editor import CRISPRBaseEditor, BaseEditResult -from models.nextgen_adcs.adc_optimizer import ADCOptimizer +from models.biologics.crispr_base_editor import CRISPRBaseEditor from models.delivery.bbb_shuttles import BBBShuttleDesigner -from drug_discovery.generation.enhanced_retrosynth import EnhancedRetrosynth +from models.nextgen_adcs.adc_optimizer import ADCOptimizer + def test_af3_docking(): docker = AlphaFold3Docking("MKTVRQERLKSIVRILERSKEPVSGAQLAEELSVSRQVIVQDIAYLRSLGYNIVAT") @@ -12,51 +16,63 @@ def test_af3_docking(): assert len(results) == 1 assert results[0].success + def test_rf_design(): designer = RFDiffusionDesigner() results = designer.design_batch(["HELI"]) assert results[0].success assert "HELI" in results[0].designed_sequence + # Similar tests for others... -@pytest.mark.parametrize("from_base,to_base", [("A","G"), ("C","T")]) + +@pytest.mark.parametrize("from_base,to_base", [("A", "G"), ("C", "T")]) def test_crispr_base_edit(from_base, to_base): editor = CRISPRBaseEditor() res = editor.base_edit("ATCG", 1, from_base, to_base) assert res.success assert res.edit_efficiency > 0.8 + def test_adc_optimize(): opt = ADCOptimizer() res = opt.optimize("payload", "Ab-shuttle") assert res.dar > 3.5 assert res.bbb_score > 0.5 + def test_bbb_shuttle(): designer = BBBShuttleDesigner() res = designer.design_shuttle("cargo") assert res.penetration_score > 0.7 + def test_enhanced_retrosynth(): synth = EnhancedRetrosynth() res = synth.plan_synthesis("CCO") assert len(res.retrosynth_paths) == 5 + def test_ai2bmd(): from drug_discovery.ai2bmd.ai2bmd_dynamics import AI2BMDDynamics + dynamics = AI2BMDDynamics() res = asyncio.run(dynamics.simulate_batch([("CCO", "pdb")])) assert res[0].success + def test_mrna_opt(): from drug_discovery.mrna_therapeutics.mrna_optimizer import mRNAOptimizer + opt = mRNAOptimizer() res = opt.optimize("antigen") assert res.expression_level > 90 + def test_evo_forecast(): from models.evolutionary_dynamics.forecast import EvolutionaryForecaster + forecaster = EvolutionaryForecaster() res = forecaster.forecast(["drug1"]) - assert res[0].resistance_time > 100 \ No newline at end of file + assert res[0].resistance_time > 100 diff --git a/tests/test_2026q2_upgrades.py b/tests/test_2026q2_upgrades.py index db0c2d6bbe..eda63a7427 100644 --- a/tests/test_2026q2_upgrades.py +++ b/tests/test_2026q2_upgrades.py @@ -1,4 +1,5 @@ """Test Suite for ZANE 2026 Q2 Upgrade Modules.""" + import numpy as np import pytest @@ -6,6 +7,7 @@ class TestScientificValidation: def test_metrics_regression(self): from drug_discovery.validation.scientific_validation import compute_metrics + y_t = np.array([1.0, 2.0, 3.0, 4.0, 5.0]) y_p = np.array([1.1, 2.1, 2.9, 4.2, 4.8]) m = compute_metrics(y_t, y_p, "regression") @@ -13,78 +15,96 @@ def test_metrics_regression(self): def test_scaffold_split(self): from drug_discovery.validation.scientific_validation import scaffold_split - smi = ["CCO","CCN","c1ccccc1","CC(=O)O","CCCC","CCC","CC=O","c1ccncc1","CC(C)O","CCCN"] + + smi = ["CCO", "CCN", "c1ccccc1", "CC(=O)O", "CCCC", "CCC", "CC=O", "c1ccncc1", "CC(C)O", "CCCN"] tr, va, te = scaffold_split(smi, 0.7, 0.15, seed=42) - assert len(tr)+len(va)+len(te) == len(smi) - assert len(set(tr)&set(va)) == 0 + assert len(tr) + len(va) + len(te) == len(smi) + assert len(set(tr) & set(va)) == 0 def test_scaffold_kfold(self): from drug_discovery.validation.scientific_validation import scaffold_kfold - folds = scaffold_kfold(["CCO","CCN","c1ccccc1","CC(=O)O","CCCC"]*4, n_folds=3) + + folds = scaffold_kfold(["CCO", "CCN", "c1ccccc1", "CC(=O)O", "CCCC"] * 4, n_folds=3) assert len(folds) == 3 def test_bootstrap_ci(self): from drug_discovery.validation.scientific_validation import bootstrap_ci - ci = bootstrap_ci(np.array([0.85,0.87,0.82,0.89,0.86]), n_bootstrap=1000) + + ci = bootstrap_ci(np.array([0.85, 0.87, 0.82, 0.89, 0.86]), n_bootstrap=1000) assert ci["ci_lower"] < ci["mean"] < ci["ci_upper"] def test_config_hash(self): from drug_discovery.validation.scientific_validation import config_hash - assert config_hash({"lr":0.001,"epochs":100}) == config_hash({"epochs":100,"lr":0.001}) + + assert config_hash({"lr": 0.001, "epochs": 100}) == config_hash({"epochs": 100, "lr": 0.001}) def test_experiment_report(self): from drug_discovery.validation.scientific_validation import ExperimentReport + r = ExperimentReport(model_name="egnn") r.fold_metrics = [{"rmse": 0.5}, {"rmse": 0.4}, {"rmse": 0.45}] r.compute_aggregates() assert r.aggregate_metrics["rmse"]["mean"] == pytest.approx(0.45, abs=0.01) + class TestDataPipeline: def test_smiles_validation(self): from drug_discovery.data.pipeline import is_valid_smiles_fast + assert is_valid_smiles_fast("CCO") and not is_valid_smiles_fast("") and not is_valid_smiles_fast("CC(") def test_validate_batch(self): from drug_discovery.data.pipeline import validate_batch - r = validate_batch(["CCO","CCN","INVALID!!!!","c1ccccc1"]) + + r = validate_batch(["CCO", "CCN", "INVALID!!!!", "c1ccccc1"]) assert r["total"] == 4 and r["valid"] >= 2 def test_lipinski(self): from drug_discovery.data.pipeline import lipinski_filter - assert lipinski_filter({"mol_weight":300,"logp":2.0,"hbd":1,"hba":3})["passes"] + + assert lipinski_filter({"mol_weight": 300, "logp": 2.0, "hbd": 1, "hba": 3})["passes"] def test_tanimoto(self): from drug_discovery.data.pipeline import tanimoto_similarity - fp = np.array([1,0,1,1,0], dtype=np.float32) + + fp = np.array([1, 0, 1, 1, 0], dtype=np.float32) assert tanimoto_similarity(fp, fp) == pytest.approx(1.0) def test_dataset(self): from drug_discovery.data.pipeline import MolecularDataset - assert MolecularDataset(smiles=["CCO","CCN"]).size == 2 + + assert MolecularDataset(smiles=["CCO", "CCN"]).size == 2 + class TestUncertainty: def test_conformal(self): from drug_discovery.evaluation.uncertainty import ConformalPredictor + cp = ConformalPredictor(alpha=0.1) - cp.calibrate(np.array([1.,2.,3.,4.,5.]), np.array([1.1,1.9,3.2,3.8,5.1])) - lo, hi = cp.predict_interval(np.array([1.,2.,3.])) + cp.calibrate(np.array([1.0, 2.0, 3.0, 4.0, 5.0]), np.array([1.1, 1.9, 3.2, 3.8, 5.1])) + lo, hi = cp.predict_interval(np.array([1.0, 2.0, 3.0])) assert all(lo < hi) def test_ece(self): from drug_discovery.evaluation.uncertainty import expected_calibration_error - assert 0 <= expected_calibration_error(np.array([.9,.8,.3,.1,.7]), np.array([1,1,0,0,1]), 5) <= 1 + + assert 0 <= expected_calibration_error(np.array([0.9, 0.8, 0.3, 0.1, 0.7]), np.array([1, 1, 0, 0, 1]), 5) <= 1 + class TestMultiObjective: def test_pareto(self): from drug_discovery.optimization.multi_objective import is_pareto_efficient - assert is_pareto_efficient(np.array([[1,4],[2,3],[3,2],[4,1],[2,2]])).sum() >= 1 + + assert is_pareto_efficient(np.array([[1, 4], [2, 3], [3, 2], [4, 1], [2, 2]])).sum() >= 1 def test_hypervolume(self): from drug_discovery.optimization.multi_objective import hypervolume_indicator - assert hypervolume_indicator(np.array([[2.,3.],[3.,2.]]), np.array([0.,0.])) > 0 + + assert hypervolume_indicator(np.array([[2.0, 3.0], [3.0, 2.0]]), np.array([0.0, 0.0])) > 0 def test_gp(self): from drug_discovery.optimization.multi_objective import GaussianProcessSurrogate + gp = GaussianProcessSurrogate() gp.fit(np.random.rand(10, 3), np.random.rand(10)) m, v = gp.predict(np.random.rand(2, 3)) @@ -92,7 +112,14 @@ def test_gp(self): def test_mobo(self): from drug_discovery.optimization.multi_objective import MOBOConfig, MultiObjectiveBayesianOptimizer - opt = MultiObjectiveBayesianOptimizer(MOBOConfig(objective_names=["a","b"], - objective_directions=["maximize","maximize"], ref_point=[0.,0.], num_mc_samples=10)) - opt.tell(np.random.rand(5,3), np.random.rand(5,2)) + + opt = MultiObjectiveBayesianOptimizer( + MOBOConfig( + objective_names=["a", "b"], + objective_directions=["maximize", "maximize"], + ref_point=[0.0, 0.0], + num_mc_samples=10, + ) + ) + opt.tell(np.random.rand(5, 3), np.random.rand(5, 2)) assert opt.summary()["total_observations"] == 5 diff --git a/tests/test_advanced_modules_comprehensive.py b/tests/test_advanced_modules_comprehensive.py index d774bc340e..ceafc8fdf5 100644 --- a/tests/test_advanced_modules_comprehensive.py +++ b/tests/test_advanced_modules_comprehensive.py @@ -24,7 +24,7 @@ def test_protein_structure_validation(self): "large": {"atoms": 10000, "chains": 5}, } - for name, struct in structures.items(): + for _name, struct in structures.items(): assert isinstance(struct, dict) def test_protein_to_graph_conversion(self): @@ -157,9 +157,7 @@ def test_kg_query(self): """Test querying KG""" with patch("drug_discovery.knowledge_graph.knowledge_graph.KnowledgeGraph") as mock_kg_class: mock_kg = MagicMock() - mock_kg.query.return_value = [ - {"subject": "drug", "predicate": "treats", "object": "disease"} - ] + mock_kg.query.return_value = [{"subject": "drug", "predicate": "treats", "object": "disease"}] mock_kg_class.return_value = mock_kg diff --git a/tests/test_compliance_api_supply.py b/tests/test_compliance_api_supply.py index c78c42469f..e7c55cfbc1 100644 --- a/tests/test_compliance_api_supply.py +++ b/tests/test_compliance_api_supply.py @@ -229,13 +229,15 @@ def save(user=None): def test_require_signature_passes(self): role = Role.from_template("lead") user = User( - user_id="u3", name="Lead", + user_id="u3", + name="Lead", role=role, password_hash=sha256_hash(None), # we set manually is_active=True, ) # Set a real password hash import hashlib + user.password_hash = hashlib.sha256(b"u3:mypass").hexdigest() @require_signature(reason="test export") diff --git a/tests/test_compliance_hardening_integration.py b/tests/test_compliance_hardening_integration.py index 5db0d0c0fd..a4fc46c9ef 100644 --- a/tests/test_compliance_hardening_integration.py +++ b/tests/test_compliance_hardening_integration.py @@ -1,18 +1,18 @@ """Test script for compliance hardening examples.""" import sys -sys.path.insert(0, '/workspaces/zane') -from drug_discovery.safety.strict_compliance_gate import ( - ComplianceLevel, - StrictComplianceGate, -) +sys.path.insert(0, "/workspaces/zane") + +from drug_discovery.compliance.audit_trail import ComplianceAuditLogger from drug_discovery.safety.parametrized_toxicity_gate import ( ParametrizedToxicityGate, ToxicityThresholdConfig, ) -from drug_discovery.compliance.audit_trail import ComplianceAuditLogger - +from drug_discovery.safety.strict_compliance_gate import ( + ComplianceLevel, + StrictComplianceGate, +) print("=" * 70) print("COMPLIANCE HARDENING - BASIC FUNCTIONAL TEST") @@ -25,8 +25,8 @@ result = gate.evaluate(smiles, toxicity_probs={"herg": 0.1}) print(f" SMILES: {smiles}") print(f" Quality Tier: {result.quality_tier.value}") -print(f" Status: PASS" if result.overall_passed else " Status: FAIL") -print(f" ✓ Strict compliance evaluation working") +print(" Status: PASS" if result.overall_passed else " Status: FAIL") +print(" ✓ Strict compliance evaluation working") # Test 2: Regulatory tier switching print("\n[TEST 2] Regulatory Tier Switching") @@ -35,31 +35,29 @@ print(f" Discovery hERG threshold: {discovery.strict_herg_threshold}") print(f" Strict hERG threshold: {strict.strict_herg_threshold}") print(f" Strict is more restrictive: {strict.strict_herg_threshold < discovery.strict_herg_threshold}") -print(f" ✓ Regulatory tier switching working") +print(" ✓ Regulatory tier switching working") # Test 3: Parametrized toxicity gate print("\n[TEST 3] Parametrized Toxicity Gate") config = ToxicityThresholdConfig.from_regulatory_tier("ind") gate = ParametrizedToxicityGate(config) result = gate.evaluate(herg_prob=0.12, ames_prob=0.08) -print(f" IND Tier - Active thresholds:") +print(" IND Tier - Active thresholds:") print(f" hERG: {result['active_thresholds']['herg']}") print(f" Ames: {result['active_thresholds']['ames']}") print(f" Evaluation passed: {result['passed']}") -print(f" ✓ Parametrized toxicity gate working") +print(" ✓ Parametrized toxicity gate working") # Test 4: Audit trail print("\n[TEST 4] Audit Trail & Compliance Logging") audit_logger = ComplianceAuditLogger() audit_logger.log_compound_screened("CCO", "COMP-001", "test_user") -audit_logger.log_quality_assessment( - "COMP-001", "CCO", "TIER_1", True, ["none"], "test_user" -) +audit_logger.log_quality_assessment("COMP-001", "CCO", "TIER_1", True, ["none"], "test_user") is_verified = audit_logger.verify_integrity() report = audit_logger.export_report() print(f" Audit entries recorded: {report['total_entries']}") print(f" Chain integrity verified: {is_verified}") -print(f" ✓ Audit trail working") +print(" ✓ Audit trail working") print("\n" + "=" * 70) print("✓ ALL COMPLIANCE HARDENING TESTS PASSED") diff --git a/tests/test_dashboard_comprehensive.py b/tests/test_dashboard_comprehensive.py index b811d96964..d032ae6821 100644 --- a/tests/test_dashboard_comprehensive.py +++ b/tests/test_dashboard_comprehensive.py @@ -55,12 +55,27 @@ def test_snapshot_creation_valid(self): def test_snapshot_creation_all_fields(self): """Test all snapshot fields are accessible""" snapshot = DashboardSnapshot( - run_id="r1", model_type="transformer", mode="screening", - molecules_screened=200, molecules_generated=100, active_jobs=3, - hit_rate=0.82, avg_qed=0.80, avg_sa=2.5, best_binding=-9.0, - epoch=5, total_epochs=50, train_loss=0.03, val_loss=0.04, - latency_ms=200.0, user_query="test", filter_query="test", - cpu_util=30.0, gpu_util=60.0, memory_gb=5.0, tick=50, + run_id="r1", + model_type="transformer", + mode="screening", + molecules_screened=200, + molecules_generated=100, + active_jobs=3, + hit_rate=0.82, + avg_qed=0.80, + avg_sa=2.5, + best_binding=-9.0, + epoch=5, + total_epochs=50, + train_loss=0.03, + val_loss=0.04, + latency_ms=200.0, + user_query="test", + filter_query="test", + cpu_util=30.0, + gpu_util=60.0, + memory_gb=5.0, + tick=50, ) assert snapshot.molecules_screened == 200 assert snapshot.train_loss == 0.03 @@ -69,12 +84,27 @@ def test_snapshot_creation_all_fields(self): def test_snapshot_edge_case_zero_values(self): """Test snapshot with zero values""" snapshot = DashboardSnapshot( - run_id="r0", model_type="gnn", mode="screening", - molecules_screened=0, molecules_generated=0, active_jobs=0, - hit_rate=0.0, avg_qed=0.0, avg_sa=0.0, best_binding=0.0, - epoch=0, total_epochs=0, train_loss=0.0, val_loss=0.0, - latency_ms=0.0, user_query="", filter_query="", - cpu_util=0.0, gpu_util=0.0, memory_gb=0.0, tick=0, + run_id="r0", + model_type="gnn", + mode="screening", + molecules_screened=0, + molecules_generated=0, + active_jobs=0, + hit_rate=0.0, + avg_qed=0.0, + avg_sa=0.0, + best_binding=0.0, + epoch=0, + total_epochs=0, + train_loss=0.0, + val_loss=0.0, + latency_ms=0.0, + user_query="", + filter_query="", + cpu_util=0.0, + gpu_util=0.0, + memory_gb=0.0, + tick=0, ) assert snapshot.molecules_screened == 0 assert snapshot.hit_rate == 0.0 @@ -82,12 +112,27 @@ def test_snapshot_edge_case_zero_values(self): def test_snapshot_edge_case_high_values(self): """Test snapshot with high values""" snapshot = DashboardSnapshot( - run_id="r_high", model_type="ensemble", mode="generation", - molecules_screened=1000000, molecules_generated=999999, active_jobs=1000, - hit_rate=1.0, avg_qed=1.0, avg_sa=10.0, best_binding=-100.0, - epoch=10000, total_epochs=10000, train_loss=999.99, val_loss=999.99, - latency_ms=99999.9, user_query="x"*1000, filter_query="x"*1000, - cpu_util=100.0, gpu_util=100.0, memory_gb=999.9, tick=999999, + run_id="r_high", + model_type="ensemble", + mode="generation", + molecules_screened=1000000, + molecules_generated=999999, + active_jobs=1000, + hit_rate=1.0, + avg_qed=1.0, + avg_sa=10.0, + best_binding=-100.0, + epoch=10000, + total_epochs=10000, + train_loss=999.99, + val_loss=999.99, + latency_ms=99999.9, + user_query="x" * 1000, + filter_query="x" * 1000, + cpu_util=100.0, + gpu_util=100.0, + memory_gb=999.9, + tick=999999, ) assert snapshot.molecules_screened == 1000000 assert snapshot.gpu_util == 100.0 @@ -126,7 +171,7 @@ def test_all_themes_exist(self): def test_theme_colors_defined(self): """Test each theme has required colors""" required_colors = ["primary", "secondary", "accent", "caution", "ok"] - for theme_name, theme in _DASHBOARD_THEMES.items(): + for _theme_name, theme in _DASHBOARD_THEMES.items(): for color_attr in required_colors: assert hasattr(theme, color_attr) assert isinstance(getattr(theme, color_attr), str) @@ -322,12 +367,27 @@ def test_advisor_summarize_returns_string(self): """Test advisor summarize returns string""" advisor = DashboardAIAdvisor(model_id=None) snapshot = DashboardSnapshot( - run_id="test", model_type="gnn", mode="generation", - molecules_screened=100, molecules_generated=50, active_jobs=2, - hit_rate=0.7, avg_qed=0.8, avg_sa=3.0, best_binding=-8.0, - epoch=5, total_epochs=10, train_loss=0.05, val_loss=0.06, - latency_ms=100.0, user_query="test", filter_query="mw<500", - cpu_util=50.0, gpu_util=75.0, memory_gb=10.0, tick=50, + run_id="test", + model_type="gnn", + mode="generation", + molecules_screened=100, + molecules_generated=50, + active_jobs=2, + hit_rate=0.7, + avg_qed=0.8, + avg_sa=3.0, + best_binding=-8.0, + epoch=5, + total_epochs=10, + train_loss=0.05, + val_loss=0.06, + latency_ms=100.0, + user_query="test", + filter_query="mw<500", + cpu_util=50.0, + gpu_util=75.0, + memory_gb=10.0, + tick=50, ) summary = advisor.summarize(snapshot) assert isinstance(summary, str) @@ -396,12 +456,27 @@ def test_theme_to_glyph_integration(self): def test_animated_bar_with_snapshot(self): """Test animated bar renders snapshot progress""" snapshot = DashboardSnapshot( - run_id="test", model_type="gnn", mode="generation", - molecules_screened=100, molecules_generated=50, active_jobs=2, - hit_rate=0.75, avg_qed=0.8, avg_sa=3.0, best_binding=-8.0, - epoch=15, total_epochs=20, train_loss=0.04, val_loss=0.05, - latency_ms=120.0, user_query="test", filter_query="mw<500", - cpu_util=50.0, gpu_util=75.0, memory_gb=10.0, tick=50, + run_id="test", + model_type="gnn", + mode="generation", + molecules_screened=100, + molecules_generated=50, + active_jobs=2, + hit_rate=0.75, + avg_qed=0.8, + avg_sa=3.0, + best_binding=-8.0, + epoch=15, + total_epochs=20, + train_loss=0.04, + val_loss=0.05, + latency_ms=120.0, + user_query="test", + filter_query="mw<500", + cpu_util=50.0, + gpu_util=75.0, + memory_gb=10.0, + tick=50, ) progress = snapshot.epoch / snapshot.total_epochs bar = _animated_bar(progress, snapshot.tick) @@ -415,12 +490,27 @@ def test_advisor_with_various_snapshot_modes(self): for mode in modes: snapshot = DashboardSnapshot( - run_id=f"test_{mode}", model_type="gnn", mode=mode, - molecules_screened=100, molecules_generated=50, active_jobs=2, - hit_rate=0.7, avg_qed=0.8, avg_sa=3.0, best_binding=-8.0, - epoch=5, total_epochs=10, train_loss=0.05, val_loss=0.06, - latency_ms=100.0, user_query="test", filter_query="mw<500", - cpu_util=50.0, gpu_util=75.0, memory_gb=10.0, tick=50, + run_id=f"test_{mode}", + model_type="gnn", + mode=mode, + molecules_screened=100, + molecules_generated=50, + active_jobs=2, + hit_rate=0.7, + avg_qed=0.8, + avg_sa=3.0, + best_binding=-8.0, + epoch=5, + total_epochs=10, + train_loss=0.05, + val_loss=0.06, + latency_ms=100.0, + user_query="test", + filter_query="mw<500", + cpu_util=50.0, + gpu_util=75.0, + memory_gb=10.0, + tick=50, ) summary = advisor.summarize(snapshot) assert isinstance(summary, str) diff --git a/tests/test_data_layer.py b/tests/test_data_layer.py index 4bb0b2d829..85751606c0 100644 --- a/tests/test_data_layer.py +++ b/tests/test_data_layer.py @@ -54,10 +54,12 @@ def test_is_valid_molecule(self): def test_normalize_dataframe(self): """Test DataFrame normalization.""" - df = pd.DataFrame({ - "smiles": self.test_smiles + ["CCO"], # Include duplicate - "activity": [1.0, 2.0, 3.0, 1.5], - }) + df = pd.DataFrame( + { + "smiles": [*self.test_smiles, "CCO"], # Include duplicate + "activity": [1.0, 2.0, 3.0, 1.5], + } + ) normalized_df = self.normalizer.normalize_dataframe( df, @@ -70,9 +72,11 @@ def test_normalize_dataframe(self): def test_apply_filters(self): """Test molecular filters.""" - df = pd.DataFrame({ - "smiles": self.test_smiles, - }) + df = pd.DataFrame( + { + "smiles": self.test_smiles, + } + ) filtered_df = self.normalizer.apply_filters( df, @@ -157,10 +161,12 @@ def teardown_method(self): def test_create_version(self): """Test dataset version creation.""" - df = pd.DataFrame({ - "smiles": ["CCO", "CC(C)O"], - "activity": [1.0, 2.0], - }) + df = pd.DataFrame( + { + "smiles": ["CCO", "CC(C)O"], + "activity": [1.0, 2.0], + } + ) version_id = self.versioning.create_version( df, @@ -173,10 +179,12 @@ def test_create_version(self): def test_load_version(self): """Test loading dataset version.""" - df = pd.DataFrame({ - "smiles": ["CCO", "CC(C)O"], - "activity": [1.0, 2.0], - }) + df = pd.DataFrame( + { + "smiles": ["CCO", "CC(C)O"], + "activity": [1.0, 2.0], + } + ) version_id = self.versioning.create_version(df, version_name="test") loaded_df = self.versioning.load_version(version_id) @@ -212,10 +220,12 @@ class TestMolecularDataset: def setup_method(self): """Setup test fixtures.""" - self.df = pd.DataFrame({ - "smiles": ["CCO", "CC(C)O", "CCCO"], - "target": [1.0, 0.0, 1.0], - }) + self.df = pd.DataFrame( + { + "smiles": ["CCO", "CC(C)O", "CCCO"], + "target": [1.0, 0.0, 1.0], + } + ) def test_fingerprint_featurization(self): """Test fingerprint featurization.""" @@ -227,7 +237,7 @@ def test_fingerprint_featurization(self): ) assert len(dataset) > 0 - feature, target = dataset[0] + feature, _target = dataset[0] assert feature.shape[0] > 0 def test_descriptor_featurization(self): @@ -240,7 +250,7 @@ def test_descriptor_featurization(self): ) assert len(dataset) > 0 - feature, target = dataset[0] + feature, _target = dataset[0] assert feature.shape[0] > 0 def test_graph_featurization(self): @@ -253,7 +263,7 @@ def test_graph_featurization(self): ) assert len(dataset) > 0 - feature, target = dataset[0] + feature, _target = dataset[0] assert isinstance(feature, dict) assert "atom_features" in feature assert "adjacency" in feature diff --git a/tests/test_data_training_comprehensive.py b/tests/test_data_training_comprehensive.py index 35d278cc60..fd27181ae5 100644 --- a/tests/test_data_training_comprehensive.py +++ b/tests/test_data_training_comprehensive.py @@ -23,13 +23,13 @@ class TestDataCollectorBasics: def test_data_collector_init(self): """Test DataCollector initialization""" - with patch.dict('os.environ', {'HOME': '/tmp'}): + with patch.dict("os.environ", {"HOME": "/tmp"}): collector = DataCollector() assert collector is not None def test_data_collector_with_cache_dir(self): """Test DataCollector with custom cache directory""" - with patch.dict('os.environ', {'HOME': '/tmp'}): + with patch.dict("os.environ", {"HOME": "/tmp"}): collector = DataCollector(cache_dir="/tmp/test_cache") assert collector is not None @@ -59,14 +59,13 @@ def test_collect_from_chembl_mock(self, mock_get): def test_collect_approved_drugs_mock(self): """Test collecting approved drugs""" - with patch.dict('os.environ', {'HOME': '/tmp'}): + with patch.dict("os.environ", {"HOME": "/tmp"}): collector = DataCollector() # Mock the actual method with patch.object(collector, "collect_approved_drugs") as mock_collect: - mock_collect.return_value = pd.DataFrame({ - "smiles": ["CC(=O)O", "CC(=O)OC"], - "name": ["acetic_acid", "methyl_acetate"] - }) + mock_collect.return_value = pd.DataFrame( + {"smiles": ["CC(=O)O", "CC(=O)OC"], "name": ["acetic_acid", "methyl_acetate"]} + ) df = collector.collect_approved_drugs() assert len(df) == 2 @@ -107,7 +106,7 @@ def test_molecular_dataset_iteration(self): dataset = MolecularDataset(smiles_list, targets) count = 0 - for item in dataset: + for _item in dataset: count += 1 assert count == 3 @@ -206,10 +205,7 @@ def test_empty_dataframe_handling(self): def test_missing_values_handling(self): """Test handling of missing values""" - df = pd.DataFrame({ - "smiles": ["CC(=O)O", None, "CC(=O)N"], - "property": [1.0, 2.0, None] - }) + df = pd.DataFrame({"smiles": ["CC(=O)O", None, "CC(=O)N"], "property": [1.0, 2.0, None]}) assert df.isnull().sum().sum() > 0 @@ -218,14 +214,8 @@ class TestDataMerging: def test_merge_two_dataframes(self): """Test merging two DataFrames""" - df1 = pd.DataFrame({ - "smiles": ["CC(=O)O", "CC(=O)OC"], - "source": ["source1", "source1"] - }) - df2 = pd.DataFrame({ - "smiles": ["CC(=O)N", "CN1C"], - "source": ["source2", "source2"] - }) + df1 = pd.DataFrame({"smiles": ["CC(=O)O", "CC(=O)OC"], "source": ["source1", "source1"]}) + df2 = pd.DataFrame({"smiles": ["CC(=O)N", "CN1C"], "source": ["source2", "source2"]}) merged = pd.concat([df1, df2], ignore_index=True) assert len(merged) == 4 @@ -257,7 +247,7 @@ class TestDataQuality: def test_data_quality_report(self): """Test data quality report generation""" - with patch.dict('os.environ', {'HOME': '/tmp'}): + with patch.dict("os.environ", {"HOME": "/tmp"}): collector = DataCollector() with patch.object(collector, "generate_data_quality_report") as mock_report: mock_report.return_value = { @@ -273,19 +263,14 @@ def test_data_quality_report(self): def test_duplicate_detection(self): """Test duplicate detection""" - df = pd.DataFrame({ - "smiles": ["CC(=O)O", "CC(=O)O", "CC(=O)N", "CC(=O)N"] - }) + df = pd.DataFrame({"smiles": ["CC(=O)O", "CC(=O)O", "CC(=O)N", "CC(=O)N"]}) duplicates = df.duplicated(subset=["smiles"]) assert duplicates.sum() == 2 def test_missing_data_detection(self): """Test missing data detection""" - df = pd.DataFrame({ - "smiles": ["CC(=O)O", None, "CC(=O)N"], - "property": [1.0, 2.0, None] - }) + df = pd.DataFrame({"smiles": ["CC(=O)O", None, "CC(=O)N"], "property": [1.0, 2.0, None]}) missing = df.isnull().sum() assert missing.sum() > 0 @@ -296,10 +281,7 @@ class TestDataSplitting: def test_random_split(self): """Test random train-test split""" - df = pd.DataFrame({ - "smiles": [f"MOL{i}" for i in range(100)], - "property": np.random.randn(100) - }) + df = pd.DataFrame({"smiles": [f"MOL{i}" for i in range(100)], "property": np.random.randn(100)}) train_size = int(0.8 * len(df)) train = df[:train_size] @@ -310,10 +292,7 @@ def test_random_split(self): def test_stratified_split(self): """Test stratified splitting""" - df = pd.DataFrame({ - "smiles": [f"MOL{i}" for i in range(100)], - "class": [0] * 70 + [1] * 30 - }) + df = pd.DataFrame({"smiles": [f"MOL{i}" for i in range(100)], "class": [0] * 70 + [1] * 30}) # Simple stratification class_0 = df[df["class"] == 0] @@ -324,10 +303,7 @@ def test_stratified_split(self): def test_scaffold_split(self): """Test scaffold-based splitting""" - df = pd.DataFrame({ - "smiles": [f"MOL{i}" for i in range(100)], - "scaffold": [i % 10 for i in range(100)] - }) + df = pd.DataFrame({"smiles": [f"MOL{i}" for i in range(100)], "scaffold": [i % 10 for i in range(100)]}) # Group by scaffold for scaffold in df["scaffold"].unique(): @@ -343,7 +319,7 @@ def test_batch_creation(self): data = list(range(100)) batch_size = 32 - batches = [data[i:i+batch_size] for i in range(0, len(data), batch_size)] + batches = [data[i : i + batch_size] for i in range(0, len(data), batch_size)] assert len(batches) == 4 assert len(batches[0]) == 32 @@ -355,7 +331,7 @@ def test_batch_iteration(self): num_batches = 0 for i in range(0, len(df), batch_size): - _batch = df[i:i+batch_size] + _batch = df[i : i + batch_size] num_batches += 1 assert num_batches == 4 @@ -377,10 +353,9 @@ class TestDataCaching: def test_cache_directory_creation(self): """Test cache directory is created""" - with patch.dict('os.environ', {'HOME': '/tmp'}): - with patch("os.makedirs"): - collector = DataCollector(cache_dir="/tmp/test_cache") - assert collector is not None + with patch.dict("os.environ", {"HOME": "/tmp"}), patch("os.makedirs"): + collector = DataCollector(cache_dir="/tmp/test_cache") + assert collector is not None def test_cache_file_handling(self): """Test cache file handling""" diff --git a/tests/test_database_drugs_trials.py b/tests/test_database_drugs_trials.py index 192fb76e9c..8cf85fa422 100644 --- a/tests/test_database_drugs_trials.py +++ b/tests/test_database_drugs_trials.py @@ -7,8 +7,6 @@ from __future__ import annotations -import tempfile -from pathlib import Path from unittest.mock import MagicMock, patch import pandas as pd @@ -16,7 +14,6 @@ from drug_discovery.data.collector import DataCollector - # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- @@ -127,9 +124,7 @@ def test_drug_conditions_annotation(self, tmp_path, collector): """CSV may carry condition/indication column that should survive.""" csv = tmp_path / "db_cond.csv" csv.write_text( - f"smiles,name,condition\n" - f"{ASPIRIN_SMILES},Aspirin,pain\n" - f"{IBUPROFEN_SMILES},Ibuprofen,inflammation\n" + f"smiles,name,condition\n" f"{ASPIRIN_SMILES},Aspirin,pain\n" f"{IBUPROFEN_SMILES},Ibuprofen,inflammation\n" ) df = collector.collect_from_drugbank(file_path=str(csv)) assert not df.empty @@ -153,11 +148,9 @@ def _make_pubchem_compound(self, smiles: str, name: str): return compound def test_returns_dataframe_on_success(self, collector): - compound = self._make_pubchem_compound(ASPIRIN_SMILES, "aspirin") + self._make_pubchem_compound(ASPIRIN_SMILES, "aspirin") with patch("drug_discovery.data.collector.DataCollector.collect_from_pubchem") as mock: - mock.return_value = pd.DataFrame( - [{"smiles": ASPIRIN_SMILES, "name": "aspirin", "source": "pubchem"}] - ) + mock.return_value = pd.DataFrame([{"smiles": ASPIRIN_SMILES, "name": "aspirin", "source": "pubchem"}]) df = collector.collect_from_pubchem(query="aspirin", limit=1) assert isinstance(df, pd.DataFrame) @@ -176,9 +169,7 @@ def test_drug_condition_queries(self, collector, condition): def test_schema_when_populated(self, collector): with patch("drug_discovery.data.collector.DataCollector.collect_from_pubchem") as mock: - mock.return_value = pd.DataFrame( - [{"smiles": CAFFEINE_SMILES, "name": "caffeine", "source": "pubchem"}] - ) + mock.return_value = pd.DataFrame([{"smiles": CAFFEINE_SMILES, "name": "caffeine", "source": "pubchem"}]) df = collector.collect_from_pubchem(query="caffeine", limit=5) if not df.empty: for col in ("smiles", "name", "source"): @@ -198,9 +189,7 @@ def _mock_chembl(self, records): mock_client = MagicMock() mock_client.activity.filter.return_value.only.return_value.__getitem__.return_value = records mock_client.molecule.filter.return_value.only.return_value.__getitem__.return_value = records - mock_client.target.filter.return_value.__getitem__ = lambda self, key: [ - {"target_chembl_id": "CHEMBL202"} - ] + mock_client.target.filter.return_value.__getitem__ = lambda self, key: [{"target_chembl_id": "CHEMBL202"}] return mock_client def test_returns_dataframe_on_success(self, collector): @@ -211,16 +200,16 @@ def test_returns_dataframe_on_success(self, collector): "pref_name": "aspirin", } ] - mock_new_client = self._mock_chembl(records) + self._mock_chembl(records) with patch("drug_discovery.data.collector.DataCollector.collect_from_chembl") as mock: - mock.return_value = pd.DataFrame( - [{"smiles": ASPIRIN_SMILES, "name": "aspirin", "source": "chembl"}] - ) + mock.return_value = pd.DataFrame([{"smiles": ASPIRIN_SMILES, "name": "aspirin", "source": "chembl"}]) df = collector.collect_from_chembl(target="kinase", limit=5) assert isinstance(df, pd.DataFrame) def test_empty_on_chembl_missing(self, collector): - with patch.dict("sys.modules", {"chembl_webresource_client": None, "chembl_webresource_client.new_client": None}): + with patch.dict( + "sys.modules", {"chembl_webresource_client": None, "chembl_webresource_client.new_client": None} + ): df = collector.collect_from_chembl(limit=5) assert isinstance(df, pd.DataFrame) @@ -233,9 +222,7 @@ def test_target_condition_queries(self, collector, target): def test_schema_when_populated(self, collector): with patch("drug_discovery.data.collector.DataCollector.collect_from_chembl") as mock: - mock.return_value = pd.DataFrame( - [{"smiles": IBUPROFEN_SMILES, "name": "ibuprofen", "source": "chembl"}] - ) + mock.return_value = pd.DataFrame([{"smiles": IBUPROFEN_SMILES, "name": "ibuprofen", "source": "chembl"}]) df = collector.collect_from_chembl(limit=5) if not df.empty: for col in ("smiles", "name", "source"): @@ -274,9 +261,8 @@ def test_search_by_query(self, collector): entry_resp = MagicMock() entry_resp.json.return_value = self._pdb_entry("2XYZ") entry_resp.raise_for_status = MagicMock() - with patch("requests.get", return_value=entry_resp): - with patch("requests.post", return_value=search_resp): - df = collector.collect_from_pdb(query="drug", limit=1) + with patch("requests.get", return_value=entry_resp), patch("requests.post", return_value=search_resp): + df = collector.collect_from_pdb(query="drug", limit=1) assert isinstance(df, pd.DataFrame) def test_empty_on_network_failure(self, collector): @@ -482,10 +468,12 @@ def _pubchem_frame(self): return pd.DataFrame([{"smiles": ASPIRIN_SMILES, "name": "aspirin", "source": "pubchem"}]) def _chembl_frame(self): - return pd.DataFrame([ - {"smiles": ASPIRIN_SMILES, "name": "acetylsalicylic acid", "source": "chembl"}, - {"smiles": IBUPROFEN_SMILES, "name": "ibuprofen", "source": "chembl"}, - ]) + return pd.DataFrame( + [ + {"smiles": ASPIRIN_SMILES, "name": "acetylsalicylic acid", "source": "chembl"}, + {"smiles": IBUPROFEN_SMILES, "name": "ibuprofen", "source": "chembl"}, + ] + ) def _drugbank_frame(self): return pd.DataFrame([{"smiles": CAFFEINE_SMILES, "name": "caffeine", "source": "drugbank"}]) @@ -496,11 +484,13 @@ def test_merge_deduplicates_smiles(self, collector): assert merged["smiles"].duplicated().sum() == 0 def test_merge_preserves_all_sources(self, collector): - merged = collector.merge_datasets([ - self._pubchem_frame(), - self._chembl_frame(), - self._drugbank_frame(), - ]) + merged = collector.merge_datasets( + [ + self._pubchem_frame(), + self._chembl_frame(), + self._drugbank_frame(), + ] + ) assert not merged.empty # SMILES from all three frames should be present smiles_set = set(merged["smiles"]) diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 54f259039d..9157fe7e66 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -118,7 +118,7 @@ def test_scaffold_kfold_reproducible(self): folds_b = murcko_scaffold_kfold_split_molecular(dataset, n_splits=3, seed=7) assert len(folds_a) == len(folds_b) - for (tr_a, te_a), (tr_b, te_b) in zip(folds_a, folds_b): + for (tr_a, te_a), (tr_b, te_b) in zip(folds_a, folds_b, strict=False): assert list(tr_a.indices) == list(tr_b.indices) assert list(te_a.indices) == list(te_b.indices) diff --git a/tests/test_drug_phases.py b/tests/test_drug_phases.py index e0acd6b5b6..480ab8b013 100644 --- a/tests/test_drug_phases.py +++ b/tests/test_drug_phases.py @@ -1,19 +1,21 @@ -import pytest from drug_discovery.data.rdkit_utils import smiles_to_mols from drug_discovery.generation.torchdrug_generator import TorchDrugGenerator from drug_discovery.screening.admet_models import ADMETScreen + def test_rdkit_utils(): smiles = ["CCO"] mols = smiles_to_mols(smiles) assert len(mols) == 1 + def test_torchdrug_generate(): gen = TorchDrugGenerator() smiles = gen.generate(10) assert len(smiles) == 10 + def test_admet_screen(): screen = ADMETScreen() preds = screen.predict(["CCO"]) - assert 'herg' in preds \ No newline at end of file + assert "herg" in preds diff --git a/tests/test_drugmaking.py b/tests/test_drugmaking.py index 26a87cff12..e5f38086e8 100644 --- a/tests/test_drugmaking.py +++ b/tests/test_drugmaking.py @@ -18,6 +18,7 @@ def test_imports(self): CustomDrugmakingModule, OptimizationConfig, ) + assert CustomDrugmakingModule is not None assert CompoundTestResult is not None assert CandidateResult is not None @@ -76,7 +77,7 @@ def test_rdkit_properties_instantiation(self): props = RDKitMolecularProperties() assert props is not None - assert hasattr(props, 'available') + assert hasattr(props, "available") def test_calculate_properties(self): """Test molecular property calculation.""" @@ -112,7 +113,7 @@ def test_physics_properties_instantiation(self): physics = PhysicsBasedProperties() assert physics is not None - assert hasattr(physics, 'rdkit_props') + assert hasattr(physics, "rdkit_props") def test_predict_binding_affinity(self): """Test binding affinity prediction.""" @@ -168,8 +169,12 @@ def test_default_config(self): config = OptimizationConfig() assert config.objective_names == [ - "potency", "selectivity", "solubility", "safety", - "synthetic_accessibility", "lipophilicity" + "potency", + "selectivity", + "solubility", + "safety", + "synthetic_accessibility", + "lipophilicity", ] assert config.num_iterations == 30 assert config.batch_size == 10 @@ -412,11 +417,11 @@ def test_test_toxicity_returns_result(self): result = module.test_toxicity("CCO") assert result is not None - assert hasattr(result, 'smiles') - assert hasattr(result, 'effectiveness') - assert hasattr(result, 'safety') - assert hasattr(result, 'molecular_properties') - assert hasattr(result, 'physics_properties') + assert hasattr(result, "smiles") + assert hasattr(result, "effectiveness") + assert hasattr(result, "safety") + assert hasattr(result, "molecular_properties") + assert hasattr(result, "physics_properties") def test_featurize_smiles(self): """Test SMILES featurization.""" @@ -497,6 +502,7 @@ def test_add_known_antidote(self): # Try adding a unique antidote (use timestamp to ensure uniqueness) import time + unique_antidote = f"TEST_ANTIDOTE_{int(time.time() * 1000)}" finder.add_known_antidote(unique_antidote) assert len(finder._known_antidotes) == initial_count + 1 diff --git a/tests/test_edge_cases_integration.py b/tests/test_edge_cases_integration.py index f046e7a862..55e9fb6e58 100644 --- a/tests/test_edge_cases_integration.py +++ b/tests/test_edge_cases_integration.py @@ -32,7 +32,7 @@ def test_invalid_model_type(self): def test_missing_dependencies(self): """Test handling of missing dependencies""" - with patch.dict('sys.modules', {'nonexistent_module': None}): + with patch.dict("sys.modules", {"nonexistent_module": None}): try: # Simulate import pass @@ -43,9 +43,9 @@ def test_corrupted_data_handling(self): """Test handling of corrupted data""" corrupted_data = [ None, - float('nan'), - float('inf'), - -float('inf'), + float("nan"), + float("inf"), + -float("inf"), ] for data in corrupted_data: @@ -178,7 +178,7 @@ def test_division_by_small_number(self): denominator = 1e-15 # Should handle without overflow - with np.errstate(over='ignore'): + with np.errstate(over="ignore"): result = numerator / np.maximum(denominator, 1e-10) assert np.isfinite(result) @@ -236,10 +236,7 @@ def test_data_loading_concurrency(self): # Simulate data loading samples_per_worker = len(data) // num_workers - loaded = [ - data[i*samples_per_worker:(i+1)*samples_per_worker] - for i in range(num_workers) - ] + loaded = [data[i * samples_per_worker : (i + 1) * samples_per_worker] for i in range(num_workers)] assert len(loaded) == num_workers @@ -262,7 +259,7 @@ def test_create_temp_directory(self): assert os.path.exists(temp_dir) # Create file in temp dir test_file = os.path.join(temp_dir, "test.txt") - with open(test_file, 'w') as f: + with open(test_file, "w") as f: f.write("test") assert os.path.exists(test_file) @@ -468,6 +465,7 @@ class TestVersioning: def test_numpy_version(self): """Test numpy version availability""" import numpy + version = numpy.__version__ assert isinstance(version, str) assert len(version) > 0 @@ -475,6 +473,7 @@ def test_numpy_version(self): def test_torch_version(self): """Test torch version availability""" import torch + version = torch.__version__ assert isinstance(version, str) assert len(version) > 0 @@ -482,6 +481,7 @@ def test_torch_version(self): def test_pandas_version(self): """Test pandas version availability""" import pandas + version = pandas.__version__ assert isinstance(version, str) assert len(version) > 0 @@ -493,16 +493,19 @@ class TestDocstringCoverage: def test_module_docstrings(self): """Test modules have docstrings""" from drug_discovery.models import ensemble + assert ensemble.__doc__ is not None or ensemble is not None def test_class_docstrings(self): """Test classes have docstrings""" from drug_discovery.models.ensemble import EnsembleModel + assert EnsembleModel.__doc__ is not None or EnsembleModel is not None def test_function_docstrings(self): """Test functions have docstrings""" from drug_discovery.dashboard import _resolve_theme + assert _resolve_theme.__doc__ is not None or _resolve_theme is not None diff --git a/tests/test_ensemble_comprehensive.py b/tests/test_ensemble_comprehensive.py index dbe11a62be..769cd8991a 100644 --- a/tests/test_ensemble_comprehensive.py +++ b/tests/test_ensemble_comprehensive.py @@ -14,6 +14,7 @@ class MockModel(nn.Module): """Simple mock model for testing""" + def __init__(self, output_dim=1): super().__init__() self.output_dim = output_dim @@ -25,6 +26,7 @@ def forward(self, x): class MockGNNModel(nn.Module): """Mock GNN model""" + def __init__(self): super().__init__() self.linear = nn.Linear(10, 32) @@ -35,6 +37,7 @@ def forward(self, x): class MockTransformerModel(nn.Module): """Mock Transformer model""" + def __init__(self): super().__init__() self.linear = nn.Linear(10, 32) diff --git a/tests/test_env_and_active_learning.py b/tests/test_env_and_active_learning.py index 50b4b8b544..edbecf469d 100644 --- a/tests/test_env_and_active_learning.py +++ b/tests/test_env_and_active_learning.py @@ -1,9 +1,10 @@ """Unit tests for environmental tests, LIMS optimizer, active learning sampler, and ABFE residuals.""" -import unittest -from infrastructure.lims.latency_optimizer import LimsLatencyOptimizer, get_default_optimizer + import importlib.util -import os import pathlib +import unittest + +from infrastructure.lims.latency_optimizer import LimsLatencyOptimizer, get_default_optimizer # Import uncertainty_sampler directly from its file to avoid package-level heavy deps _this_dir = pathlib.Path(__file__).resolve().parent @@ -12,17 +13,19 @@ uncertainty_mod = importlib.util.module_from_spec(spec) spec.loader.exec_module(uncertainty_mod) UncertaintySampler = uncertainty_mod.UncertaintySampler -from drug_discovery.smd.abfe_residuals import compute_residuals, rmse, summarize_abfe from drug_discovery.safety.environmental_tests import run_environmental_tests +from drug_discovery.smd.abfe_residuals import compute_residuals, summarize_abfe class TestEnvAndAL(unittest.TestCase): def test_lims_optimizer_basic(self): opt = LimsLatencyOptimizer(cache_ttl=1.0) + # simple function def f(x): return x * 2 - wrapped = opt.instrument(lambda x: str(x))(f) # key_func passed wrongly on purpose -> fallback + + opt.instrument(lambda x: str(x))(f) # key_func passed wrongly on purpose -> fallback # call should succeed self.assertEqual(f(3), 6) # default optimizer exists @@ -30,12 +33,12 @@ def f(x): def test_uncertainty_sampler(self): sampler = UncertaintySampler() - smiles = ['A', 'B', 'C', 'D'] + smiles = ["A", "B", "C", "D"] uncertainties = [0.1, 0.9, 0.4, 0.8] selected = sampler.select_batch(smiles, uncertainties, batch_size=2) self.assertEqual(len(selected), 2) # highest uncertainty items should be included - self.assertTrue('B' in selected or 'D' in selected) + self.assertTrue("B" in selected or "D" in selected) def test_abfe_residuals(self): pred = [1.0, 2.0, 3.0] @@ -43,15 +46,15 @@ def test_abfe_residuals(self): res = compute_residuals(pred, obs) self.assertEqual(len(res), 3) s = summarize_abfe(pred, obs) - self.assertIn('rmse', s) - self.assertGreaterEqual(s['rmse'], 0.0) + self.assertIn("rmse", s) + self.assertGreaterEqual(s["rmse"], 0.0) def test_environmental_tests(self): - out = run_environmental_tests('CC(=O)O') - self.assertIn('ph_profiles', out) - self.assertIn('plasma_binding', out) - self.assertIn('7.4', out['ph_profiles']) + out = run_environmental_tests("CC(=O)O") + self.assertIn("ph_profiles", out) + self.assertIn("plasma_binding", out) + self.assertIn("7.4", out["ph_profiles"]) -if __name__ == '__main__': +if __name__ == "__main__": unittest.main() diff --git a/tests/test_herg_predictor.py b/tests/test_herg_predictor.py index c307efb046..edff55a1f9 100644 --- a/tests/test_herg_predictor.py +++ b/tests/test_herg_predictor.py @@ -41,7 +41,7 @@ def test_custom_coefficients(self): def test_prediction_output_structure(self): """Test that prediction returns complete HERGPrediction object.""" result = self.predictor.predict("CCO") - + self.assertIsInstance(result, HERGPrediction) self.assertEqual(result.smiles, "CCO") self.assertIsInstance(result.inhibition_probability, float) @@ -64,11 +64,11 @@ def test_probability_bounds(self): def test_ic50_estimation(self): """Test that IC50 is estimated in realistic nanomolar range.""" result = self.predictor.predict("c1ccc(Nc2nccc(Oc3ccccc3)n2)cc1") # Example ARB-like - + # IC50 should be in plausible range (100 nM to 100 µM) self.assertGreaterEqual(result.ic50_estimate_nM, 100.0) self.assertLess(result.ic50_estimate_nM, 100000.0) - + # Range should be plausible self.assertLess(result.ic50_range_nM[0], result.ic50_range_nM[1]) @@ -77,7 +77,7 @@ def test_cipa_risk_classification(self): # Very lipophilic aromatic compound with basic N (should be high risk) high_risk = self.predictor.predict("CN1CCC[C@H]1c2ccccc2Cl") # Piperidine-phenyl self.assertIn(high_risk.cipa_risk_category, ["category_2", "category_3"]) - + # Simple polar compound should be lower risk low_risk = self.predictor.predict("CCO") # Ethanol # Ethanol should be low or category_2 (not high) @@ -87,11 +87,11 @@ def test_qtc_risk_assessment(self): """Test QTc prolongation risk is derived from hERG and potency.""" result = self.predictor.predict("CCO") self.assertIn(result.qtichan_risk, ["very_low", "low", "moderate", "high"]) - + # Risk should increase with inhibition probability high_inhibit = self.predictor.predict("CN1CCN(c2ccc(cc2)Cl)CC1") # More aromatic/basic low_inhibit = self.predictor.predict("CO") - + # This is not strictly guaranteed due to heuristics, but generally true # Just verify that risk categories are valid self.assertIn(high_inhibit.qtichan_risk, ["very_low", "low", "moderate", "high"]) @@ -108,7 +108,7 @@ def test_key_concerns_identification(self): # Molecule with high logP and low TPSA should have concerns result = self.predictor.predict("c1ccc(cc1)C(c2ccccc2)C(c3ccccc3)c4ccccc4") # Very lipophilic self.assertGreater(len(result.key_concerns), 0) - + # Simple molecule might have fewer or no concerns result_simple = self.predictor.predict("C") # Should be a list (may be empty) @@ -135,7 +135,7 @@ def test_glp_backward_compatibility(self): """Test that GLPToxPanel still produces HERGResult with correct interface.""" panel = PreClinicalToxPanel() result = panel.evaluate("CCO") - + # HERGResult should have expected attributes self.assertIsNotNone(result.herg) self.assertGreaterEqual(result.herg.inhibition_probability, 0.0) @@ -150,7 +150,7 @@ def test_herg_threshold_enforcement(self): # Very strict threshold strict_panel = PreClinicalToxPanel(herg_threshold=0.1) result = strict_panel.evaluate("c1ccccc1") # Benzene (somewhat aromatic) - + # Should fail if inhibition > 0.1 if result.herg.inhibition_probability > 0.1: self.assertFalse(result.herg.passed) @@ -161,7 +161,7 @@ def test_custom_herg_predictor_in_glp(self): """Test that custom HERGPredictor can be passed to PreClinicalToxPanel.""" custom_predictor = HERGPredictor(logp_coeff=0.3, tpsa_coeff=-0.2) panel = PreClinicalToxPanel(herg_predictor=custom_predictor) - + self.assertIs(panel.herg_predictor, custom_predictor) self.assertEqual(panel.herg_predictor.logp_coeff, 0.3) @@ -170,7 +170,7 @@ def test_batch_evaluation_uses_herg_predictor(self): panel = PreClinicalToxPanel() smiles_list = ["CCO", "c1ccccc1", "CC(=O)O"] results = panel.evaluate_batch(smiles_list) - + self.assertEqual(len(results), 3) for result in results: self.assertIsNotNone(result.herg) diff --git a/tests/test_models.py b/tests/test_models.py index a93305ff65..10fb88606c 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -5,9 +5,10 @@ import pytest torch = pytest.importorskip("torch") -from drug_discovery.models import EnsembleModel, MolecularGNN, MolecularTransformer from torch_geometric.data import Data +from drug_discovery.models import EnsembleModel, MolecularGNN, MolecularTransformer + class TestMolecularGNN: """Test Graph Neural Network model""" diff --git a/tests/test_nvidia_llm_finetune.py b/tests/test_nvidia_llm_finetune.py index 680991865f..bcd2838a27 100644 --- a/tests/test_nvidia_llm_finetune.py +++ b/tests/test_nvidia_llm_finetune.py @@ -106,4 +106,4 @@ def empty_merge(datasets): df = collect_public_molecule_data(collector=collector, sources=["pubchem"], fallback_to_approved_drugs=True) assert not df.empty - assert collector.approved_called is True \ No newline at end of file + assert collector.approved_called is True diff --git a/tests/test_physics_oracle_bayesian_gflownet.py b/tests/test_physics_oracle_bayesian_gflownet.py index d09d115d44..5ed87ffe4b 100644 --- a/tests/test_physics_oracle_bayesian_gflownet.py +++ b/tests/test_physics_oracle_bayesian_gflownet.py @@ -280,9 +280,7 @@ def test_min_reward_floor(self): assert fn({"atoms": torch.tensor([0]), "num_atoms": 1}) >= cfg.min_reward def test_default_atoms_to_smiles(self): - s = PhysicsRewardFunction._default_atoms_to_smiles( - {"atoms": torch.tensor([0, 1, 2]), "num_atoms": 3} - ) + s = PhysicsRewardFunction._default_atoms_to_smiles({"atoms": torch.tensor([0, 1, 2]), "num_atoms": 3}) assert isinstance(s, str) assert len(s) > 0 diff --git a/tests/test_pipeline_models_comprehensive.py b/tests/test_pipeline_models_comprehensive.py index 4d48c7f429..a8b42e66ff 100644 --- a/tests/test_pipeline_models_comprehensive.py +++ b/tests/test_pipeline_models_comprehensive.py @@ -62,22 +62,12 @@ class TestMolecularTransformerBasics: def test_transformer_init(self): """Test Transformer initialization""" - transformer = MolecularTransformer( - input_dim=64, - model_dim=256, - num_heads=8, - num_layers=4 - ) + transformer = MolecularTransformer(input_dim=64, model_dim=256, num_heads=8, num_layers=4) assert transformer is not None def test_transformer_forward_pass(self): """Test Transformer forward pass""" - transformer = MolecularTransformer( - input_dim=32, - model_dim=128, - num_heads=4, - num_layers=2 - ) + transformer = MolecularTransformer(input_dim=32, model_dim=128, num_heads=4, num_layers=2) x = torch.randn(4, 10, 32) # (batch, seq_len, input_dim) try: @@ -90,34 +80,19 @@ def test_transformer_forward_pass(self): def test_transformer_number_of_heads(self): """Test Transformer with different attention heads""" for num_heads in [1, 2, 4, 8]: - transformer = MolecularTransformer( - input_dim=64, - model_dim=256, - num_heads=num_heads, - num_layers=2 - ) + transformer = MolecularTransformer(input_dim=64, model_dim=256, num_heads=num_heads, num_layers=2) assert transformer is not None def test_transformer_number_of_layers(self): """Test Transformer with different layer counts""" for num_layers in [1, 2, 3, 6, 12]: - transformer = MolecularTransformer( - input_dim=64, - model_dim=256, - num_heads=8, - num_layers=num_layers - ) + transformer = MolecularTransformer(input_dim=64, model_dim=256, num_heads=8, num_layers=num_layers) assert transformer is not None def test_transformer_model_dimensions(self): """Test Transformer with different model dimensions""" for model_dim in [64, 128, 256, 512, 1024]: - transformer = MolecularTransformer( - input_dim=32, - model_dim=model_dim, - num_heads=8, - num_layers=2 - ) + transformer = MolecularTransformer(input_dim=32, model_dim=model_dim, num_heads=8, num_layers=2) assert transformer is not None @@ -127,6 +102,7 @@ class TestDrugModelingBasics: def test_drug_modeling_init(self): """Test drug modeling module initialization""" from drug_discovery.models import drug_modeling + assert drug_modeling is not None @patch("drug_discovery.models.drug_modeling.DrugModel") @@ -241,6 +217,7 @@ class TestOptimizationBasics: def test_optimization_module_imports(self): """Test optimization module can be imported""" from drug_discovery.optimization import bayesian, multi_objective + assert bayesian is not None assert multi_objective is not None @@ -304,10 +281,7 @@ class TestPipelineWorkflow: def test_pipeline_data_preparation(self): """Test pipeline data preparation stage""" - data = pd.DataFrame({ - "smiles": ["CC(=O)O", "CC(=O)OC", "CC(=O)N"], - "property": [1.0, 2.0, 3.0] - }) + data = pd.DataFrame({"smiles": ["CC(=O)O", "CC(=O)OC", "CC(=O)N"], "property": [1.0, 2.0, 3.0]}) assert len(data) == 3 assert "smiles" in data.columns @@ -336,7 +310,7 @@ def test_pipeline_training_phase(self): x = torch.randn(32, 10) y = torch.randn(32, 1) - for epoch in range(3): + for _epoch in range(3): optimizer.zero_grad() output = model(x) loss = criterion(output, y) @@ -389,7 +363,7 @@ def test_model_state_loading(self): model2.load_state_dict(state_dict) # Models should have same parameters - for p1, p2 in zip(model1.parameters(), model2.parameters()): + for p1, p2 in zip(model1.parameters(), model2.parameters(), strict=False): assert torch.allclose(p1, p2) def test_optimizer_state_dict(self): @@ -474,7 +448,7 @@ def test_num_epochs_selection(self): model = nn.Linear(10, 1) optimizer = torch.optim.SGD(model.parameters()) - for epoch in range(10): + for _epoch in range(10): x = torch.randn(4, 10) y = torch.randn(4, 1) optimizer.zero_grad() @@ -494,7 +468,7 @@ def test_early_stopping_on_plateau(self): losses = [1.0, 0.9, 0.85, 0.83, 0.82, 0.82, 0.82, 0.82] patience = 3 - best_loss = float('inf') + best_loss = float("inf") patience_count = 0 for loss in losses: @@ -514,5 +488,5 @@ def test_validation_monitoring(self): train_losses = [1.0, 0.9, 0.8, 0.7, 0.6] val_losses = [1.1, 1.0, 0.95, 0.92, 0.91] - for train_loss, val_loss in zip(train_losses, val_losses): + for train_loss, val_loss in zip(train_losses, val_losses, strict=False): assert val_loss >= train_loss * 0.9 # Validation typically higher diff --git a/tests/test_predictor_comprehensive.py b/tests/test_predictor_comprehensive.py index 96090966f3..62cfe79597 100644 --- a/tests/test_predictor_comprehensive.py +++ b/tests/test_predictor_comprehensive.py @@ -15,6 +15,7 @@ class MockPyTorchModel(torch.nn.Module): """Mock PyTorch model for testing""" + def __init__(self, input_dim=10, output_dim=1): super().__init__() self.linear = torch.nn.Linear(input_dim, output_dim) @@ -221,7 +222,7 @@ def test_calculate_lipinski_properties_all_numeric(self): props = predictor.calculate_lipinski_properties(smiles) - for key, value in props.items(): + for _key, value in props.items(): assert isinstance(value, (int, float)) def test_calculate_lipinski_properties_non_negative(self): @@ -486,6 +487,6 @@ def test_batch_predictions(self): batch = [predictor.calculate_qed(s) for s in smiles_list] # Should be identical - for ind, bat in zip(individual, batch): + for ind, bat in zip(individual, batch, strict=False): if ind is not None and bat is not None: assert abs(ind - bat) < 1e-6 diff --git a/tests/test_scraper_comprehensive.py b/tests/test_scraper_comprehensive.py index aacaae013a..087fcf1adf 100644 --- a/tests/test_scraper_comprehensive.py +++ b/tests/test_scraper_comprehensive.py @@ -49,11 +49,7 @@ class TestPubMedSearch: def test_search_basic(self, mock_get): """Test basic search query""" mock_response = MagicMock() - mock_response.json.return_value = { - "esearchresult": { - "idlist": ["12345", "67890"] - } - } + mock_response.json.return_value = {"esearchresult": {"idlist": ["12345", "67890"]}} mock_get.return_value = mock_response api = PubMedAPI() @@ -345,9 +341,7 @@ def test_search_and_fetch_pipeline(self, mock_get): """Test search followed by fetch""" # Setup mock responses search_response = MagicMock() - search_response.json.return_value = { - "esearchresult": {"idlist": ["123", "456"]} - } + search_response.json.return_value = {"esearchresult": {"idlist": ["123", "456"]}} fetch_response = MagicMock() fetch_response.text = "" diff --git a/tests/test_strict_compliance_gate.py b/tests/test_strict_compliance_gate.py index 7c338c01d1..d13663737a 100644 --- a/tests/test_strict_compliance_gate.py +++ b/tests/test_strict_compliance_gate.py @@ -1,7 +1,7 @@ """Tests for strict compliance and parametrized quality gates. Comprehensive test suite covering: -- Strict compliance evaluation +- Strict compliance evaluation - Quality tier classification - Data integrity verification - Risk factor identification @@ -11,7 +11,6 @@ """ import unittest -from datetime import datetime from drug_discovery.safety.parametrized_toxicity_gate import ( ParametrizedToxicityGate, @@ -20,7 +19,6 @@ from drug_discovery.safety.strict_compliance_gate import ( ComplianceLevel, QualityTier, - RiskFactor, StrictComplianceGate, evaluate_batch_with_strict_compliance, ) @@ -41,7 +39,7 @@ def test_config_validation(self): """Test invalid thresholds raise ValueError.""" with self.assertRaises(ValueError): ToxicityThresholdConfig(herg_threshold=1.5) - + with self.assertRaises(ValueError): ToxicityThresholdConfig(logp_min=10, logp_max=5) @@ -245,7 +243,7 @@ def test_hardened_compliance_stricter_than_relaxed(self): compliance_level=ComplianceLevel.HARDENED, strict_herg_threshold=0.1, ) - + # Hardened should have lower thresholds self.assertLess(hardened.strict_herg_threshold, relaxed.strict_herg_threshold) @@ -262,7 +260,8 @@ def test_toxicity_prob_integration(self): ) # Should have toxicity-related checks toxicity_checks = [ - c for c in assessment.compliance_checks + c + for c in assessment.compliance_checks if any(x in c.check_name for x in ["hERG", "Ames", "Hepatotoxicity"]) ] self.assertGreater(len(toxicity_checks), 0) @@ -292,7 +291,7 @@ def test_batch_evaluation(self): smiles_list, compliance_level=ComplianceLevel.RELAXED, ) - + self.assertEqual(result["total_evaluated"], 3) self.assertEqual(result["passed"] + result["rejected"], 3) self.assertGreaterEqual(result["pass_rate"], 0) @@ -302,7 +301,7 @@ def test_batch_includes_all_assessments(self): """Test batch result includes all individual assessments.""" smiles_list = ["CC(=O)O", "CC(=O)N"] result = evaluate_batch_with_strict_compliance(smiles_list) - + self.assertEqual(len(result["assessments"]), 2) for smiles in smiles_list: self.assertIn(smiles, result["assessments"]) @@ -311,7 +310,7 @@ def test_batch_critical_issues_tracked(self): """Test batch evaluation tracks critical issues.""" smiles_list = ["CC(=O)O", "invalid"] # One valid, one invalid result = evaluate_batch_with_strict_compliance(smiles_list) - + # Should have at least one critical issue (invalid SMILES) self.assertGreater(len(result["critical_issues"]), 0) @@ -323,7 +322,7 @@ def test_relaxed_vs_strict_herg_thresholds(self): """Test compliance level affects hERG thresholds.""" relaxed = StrictComplianceGate(compliance_level=ComplianceLevel.RELAXED) strict = StrictComplianceGate(compliance_level=ComplianceLevel.STRICT) - + # Strict should have lower (more restrictive) hERG threshold self.assertLess( strict.strict_herg_threshold, @@ -334,7 +333,7 @@ def test_hardened_is_most_restrictive(self): """Test hardened compliance level is most restrictive.""" relaxed = StrictComplianceGate(compliance_level=ComplianceLevel.RELAXED) hardened = StrictComplianceGate(compliance_level=ComplianceLevel.HARDENED) - + # Hardened should have strictest threshold self.assertLess( hardened.strict_herg_threshold, @@ -354,7 +353,7 @@ def test_simple_property_heuristic(self): gate = StrictComplianceGate() # Simple SMILES: acetic acid props = gate._calculate_properties("CC(=O)O") - + self.assertIn("mw", props) self.assertIn("logp", props) self.assertIn("tpsa", props) @@ -364,6 +363,7 @@ def test_simple_property_heuristic(self): def run_test_log_output() -> None: """Run tests with logging.""" import logging + logging.basicConfig(level=logging.DEBUG) unittest.main() diff --git a/tests/test_swissadme_proxy.py b/tests/test_swissadme_proxy.py index a5ffc48818..7bbf47314a 100644 --- a/tests/test_swissadme_proxy.py +++ b/tests/test_swissadme_proxy.py @@ -26,4 +26,4 @@ def test_predict_rejects_empty_smiles(self): if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/test_synthesis_comprehensive.py b/tests/test_synthesis_comprehensive.py index 81efffc409..1f9371a358 100644 --- a/tests/test_synthesis_comprehensive.py +++ b/tests/test_synthesis_comprehensive.py @@ -47,6 +47,7 @@ def test_retrosynthesis_module_attributes(self): """Test retrosynthesis module has expected attributes""" # Module should exist and be importable import drug_discovery.synthesis.retrosynthesis as ret_module + assert ret_module is not None @@ -60,6 +61,7 @@ def test_reaction_prediction_module_imports(self): def test_reaction_prediction_classes_exist(self): """Test reaction prediction has classes""" import drug_discovery.synthesis.reaction_prediction as rp_module + # Module should be importable assert rp_module is not None @@ -101,6 +103,7 @@ def test_backends_module_imports(self): def test_backends_module_attributes(self): """Test backends module structure""" import drug_discovery.synthesis.backends as backends_module + assert backends_module is not None @patch("drug_discovery.synthesis.backends.Backend") @@ -113,6 +116,7 @@ def test_multiple_backends_available(self): """Test multiple synthesis backends are available""" # Should be able to import backends import drug_discovery.synthesis.backends as backends_module + assert backends_module is not None @@ -126,6 +130,7 @@ def test_pistachio_module_imports(self): def test_pistachio_dataset_classes(self): """Test Pistachio dataset classes exist""" import drug_discovery.synthesis.pistachio_datasets as pd_module + assert pd_module is not None @patch("drug_discovery.synthesis.pistachio_datasets.PistachioDataset") diff --git a/tests/test_testing_layer.py b/tests/test_testing_layer.py index fd2f8e9024..39f50e3737 100644 --- a/tests/test_testing_layer.py +++ b/tests/test_testing_layer.py @@ -237,6 +237,7 @@ def test_compute_prediction_confidence(self): def test_batch_uncertainty_estimation(self): """Test batch uncertainty estimation.""" + def dummy_model1(smiles): return 0.5 diff --git a/tests/test_titan_architecture.py b/tests/test_titan_architecture.py index 7bb89e3416..f071c02078 100644 --- a/tests/test_titan_architecture.py +++ b/tests/test_titan_architecture.py @@ -34,9 +34,9 @@ def test_off_target_veto() -> None: scorer.score_off_targets(verapamil) veto = exc_info.value - assert veto.score > veto.threshold, ( - f"Expected veto.score ({veto.score:.3f}) > veto.threshold ({veto.threshold:.3f})" - ) + assert ( + veto.score > veto.threshold + ), f"Expected veto.score ({veto.score:.3f}) > veto.threshold ({veto.threshold:.3f})" assert veto.smiles == verapamil @@ -84,9 +84,7 @@ def test_onnx_latency(monkeypatch: pytest.MonkeyPatch) -> None: result = server.predict(smiles) elapsed_s = time.perf_counter() - t_start - assert elapsed_s < 0.15, ( - f"Inference took {elapsed_s * 1000:.1f} ms, expected < 150 ms" - ) + assert elapsed_s < 0.15, f"Inference took {elapsed_s * 1000:.1f} ms, expected < 150 ms" assert "mean_score" in result assert "variance" in result assert "confidence_warning" in result diff --git a/torch_geometric_local/__init__.py b/torch_geometric_local/__init__.py index 08d3c02c5f..5048c0e49f 100644 --- a/torch_geometric_local/__init__.py +++ b/torch_geometric_local/__init__.py @@ -11,4 +11,4 @@ from .loader import DataLoader from .nn import GATConv, MessagePassing, global_max_pool, global_mean_pool -__all__ = ["Data", "DataLoader", "GATConv", "MessagePassing", "global_mean_pool", "global_max_pool"] +__all__ = ["Data", "DataLoader", "GATConv", "MessagePassing", "global_max_pool", "global_mean_pool"] diff --git a/training/de_novo_enforcer.py b/training/de_novo_enforcer.py index 7228a21ba3..c4d742378a 100644 --- a/training/de_novo_enforcer.py +++ b/training/de_novo_enforcer.py @@ -1,16 +1,17 @@ -from rdkit import Chem -from rdkit import DataStructs -from rdkit.Chem import AllChem -import numpy as np import logging import os +from rdkit import Chem, DataStructs +from rdkit.Chem import AllChem + logger = logging.getLogger(__name__) + class DeNovoStrictEnforcer: """ Ensures that generated molecules are novel and not regurgitations of known drugs. """ + def __init__(self, threshold: float = 0.45): self.threshold = threshold self.known_fingerprints = [] @@ -27,7 +28,7 @@ def load_public_archives(self, chembl_db_path: str): logger.info(f"Loading public drug archives from {chembl_db_path}...") try: - with open(chembl_db_path, 'r') as f: + with open(chembl_db_path) as f: for line in f: smiles = line.strip().split()[0] mol = Chem.MolFromSmiles(smiles) @@ -37,7 +38,7 @@ def load_public_archives(self, chembl_db_path: str): self.known_smiles.add(smiles) logger.info(f"Loaded {len(self.known_fingerprints)} molecules into novelty enforcer.") except Exception as e: - logger.error(f"Failed to load archives: {str(e)}") + logger.error(f"Failed to load archives: {e!s}") def calculate_novelty_penalty(self, generated_smiles: str) -> float: """ @@ -46,7 +47,7 @@ def calculate_novelty_penalty(self, generated_smiles: str) -> float: """ mol = Chem.MolFromSmiles(generated_smiles) if not mol: - return -10.0 # Invalid molecule penalty + return -10.0 # Invalid molecule penalty # Check for exact match first if generated_smiles in self.known_smiles: @@ -54,7 +55,7 @@ def calculate_novelty_penalty(self, generated_smiles: str) -> float: return -1000.0 fp = AllChem.GetMorganFingerprintAsBitVect(mol, 2, nBits=2048) - + if not self.known_fingerprints: return 0.0 @@ -64,8 +65,10 @@ def calculate_novelty_penalty(self, generated_smiles: str) -> float: if max_sim > self.threshold: # Catastrophic negative reward for minor tweaks of existing drugs - penalty = -100.0 * (max_sim / self.threshold)**2 - logger.info(f"Novelty Veto: Max similarity {max_sim:.4f} exceeds threshold {self.threshold}. Penalty: {penalty:.2f}") + penalty = -100.0 * (max_sim / self.threshold) ** 2 + logger.info( + f"Novelty Veto: Max similarity {max_sim:.4f} exceeds threshold {self.threshold}. Penalty: {penalty:.2f}" + ) return penalty return 0.0 diff --git a/training/global_reward_orchestrator.py b/training/global_reward_orchestrator.py index 2131740498..93ad2afb17 100644 --- a/training/global_reward_orchestrator.py +++ b/training/global_reward_orchestrator.py @@ -1,22 +1,23 @@ -import torch import logging -import asyncio -from typing import Dict, Any, Optional +from typing import Any + +from clinical.digital_twin.epigenetic_profiler import EpigeneticAgeCalculator +from clinical.digital_twin.microbiome_metabolomics import MicrobiomeToxicityVeto, PharmacobiomiomicEngine +from models.structural.ph_dependent_protonation import MicroenvironmentIonizationEngine +from training.de_novo_enforcer import DeNovoStrictEnforcer # Import previously built modules from training.n1_health_condition_optimizer import ConditionAdaptiveRewardFunction -from training.de_novo_enforcer import DeNovoStrictEnforcer -from models.structural.ph_dependent_protonation import MicroenvironmentIonizationEngine -from clinical.digital_twin.microbiome_metabolomics import PharmacobiomiomicEngine, MicrobiomeToxicityVeto -from clinical.digital_twin.epigenetic_profiler import EpigeneticAgeCalculator logger = logging.getLogger(__name__) + class PanArchitectureReward: """ - The master orchestrator that integrates every ZANE subsystem into a + The master orchestrator that integrates every ZANE subsystem into a single unified reward signal for the Reinforcement Learning agent. """ + def __init__(self, patient_state: Any): self.patient_state = patient_state self.n1_optimizer = ConditionAdaptiveRewardFunction() @@ -24,7 +25,7 @@ def __init__(self, patient_state: Any): self.ph_engine = MicroenvironmentIonizationEngine() self.microbiome_engine = PharmacobiomiomicEngine() self.age_calculator = EpigeneticAgeCalculator() - + # Weights for the master equation self.weights = { "docking_score": 0.30, @@ -32,10 +33,10 @@ def __init__(self, patient_state: Any): "admet_safety": 0.15, "n1_metabolic_fit": 0.15, "physicochemical": 0.10, - "solubility": 0.10 + "solubility": 0.10, } - async def calculate_total_reward(self, smiles: str, predicted_properties: Dict[str, float]) -> float: + async def calculate_total_reward(self, smiles: str, predicted_properties: dict[str, float]) -> float: """ Executes asynchronous calls to all subsystems to calculate the unified reward. R_total = sum(w_i * r_i) @@ -44,27 +45,25 @@ async def calculate_total_reward(self, smiles: str, predicted_properties: Dict[s # 1. Novelty Check (Banning memorized drugs) novelty_penalty = self.novelty_enforcer.calculate_novelty_penalty(smiles) if novelty_penalty < -500: - return novelty_penalty # Immediate rejection for memorized drugs + return novelty_penalty # Immediate rejection for memorized drugs # 2. Patient-Specific Metabolic Fit (eGFR/AST/ALT) n1_reward = self.n1_optimizer.compute_n1_optimized_reward( - base_reward=0, - patient_state=self.patient_state, - predicted_properties=predicted_properties + base_reward=0, patient_state=self.patient_state, predicted_properties=predicted_properties ) # 3. Microbiome Toxicity Veto try: # Assuming patient microbiome profile is part of patient_state - microbiome_profile = getattr(self.patient_state, 'microbiome_profile', {"Bacteroides": 0.4}) + microbiome_profile = getattr(self.patient_state, "microbiome_profile", {"Bacteroides": 0.4}) self.microbiome_engine.predict_microbial_cleavage(smiles, microbiome_profile) microbiome_reward = 1.0 except MicrobiomeToxicityVeto as e: - logger.warning(f"Microbiome Veto for {smiles}: {str(e)}") + logger.warning(f"Microbiome Veto for {smiles}: {e!s}") microbiome_reward = -100.0 # 4. pH-Dependent Solubility (Localized microenvironment) - target_ph = getattr(self.patient_state, 'target_tissue_ph', 6.5) # e.g. Tumor pH + target_ph = getattr(self.patient_state, "target_tissue_ph", 6.5) # e.g. Tumor pH solubility_score = self.ph_engine.calculate_ph_dependent_solubility(smiles, target_ph) # 5. Physicochemical Constraints (Learned via RAG) @@ -74,17 +73,17 @@ async def calculate_total_reward(self, smiles: str, predicted_properties: Dict[s # Master Equation Integration total_reward = ( - self.weights["docking_score"] * base_docking + - self.weights["novelty"] * (1.0 + novelty_penalty/100.0) + - self.weights["admet_safety"] * admet_score + - self.weights["n1_metabolic_fit"] * (1.0 + n1_reward/100.0) + - self.weights["solubility"] * solubility_score + - (0.1 * microbiome_reward) # Extra weight for microbiome safety + self.weights["docking_score"] * base_docking + + self.weights["novelty"] * (1.0 + novelty_penalty / 100.0) + + self.weights["admet_safety"] * admet_score + + self.weights["n1_metabolic_fit"] * (1.0 + n1_reward / 100.0) + + self.weights["solubility"] * solubility_score + + (0.1 * microbiome_reward) # Extra weight for microbiome safety ) logger.info(f"Unified Reward for {smiles}: {total_reward:.4f}") return float(total_reward) except Exception as e: - logger.error(f"Error in reward orchestration: {str(e)}") + logger.error(f"Error in reward orchestration: {e!s}") return -10.0 diff --git a/training/n1_health_condition_optimizer.py b/training/n1_health_condition_optimizer.py index 80512165a3..3f2f3ba467 100644 --- a/training/n1_health_condition_optimizer.py +++ b/training/n1_health_condition_optimizer.py @@ -1,33 +1,32 @@ -import torch -import numpy as np -import scipy.optimize as opt -from typing import Any, Dict, Optional import logging +from typing import Any logger = logging.getLogger(__name__) + class ConditionAdaptiveRewardFunction: """ - Dynamically adjusts the reinforcement learning reward landscape based on + Dynamically adjusts the reinforcement learning reward landscape based on individual patient physiological constraints. """ + def __init__(self, renal_threshold: float = 30.0, hepatic_threshold: float = 120.0): self.renal_threshold = renal_threshold self.hepatic_threshold = hepatic_threshold - def dynamic_clearance_penalty(self, patient_state: Any, predicted_properties: Dict[str, float]) -> float: + def dynamic_clearance_penalty(self, patient_state: Any, predicted_properties: dict[str, float]) -> float: """ - Injects a penalty if a molecule's predicted clearance route conflicts + Injects a penalty if a molecule's predicted clearance route conflicts with the patient's organ impairment. - + If patient_state.eGFR < 30, molecules with high renal clearance are heavily penalized. """ penalty = 0.0 - egfr = getattr(patient_state, 'egfr', 90.0) - + egfr = getattr(patient_state, "egfr", 90.0) + # Predicted renal clearance (normalized 0-1) - renal_clearance = predicted_properties.get('renal_clearance', 0.5) - + renal_clearance = predicted_properties.get("renal_clearance", 0.5) + if egfr < self.renal_threshold: # Severe renal impairment: patient cannot clear drugs via kidneys if renal_clearance > 0.2: @@ -35,35 +34,34 @@ def dynamic_clearance_penalty(self, patient_state: Any, predicted_properties: Di impairment_factor = (self.renal_threshold - egfr) / self.renal_threshold penalty -= 100.0 * impairment_factor * renal_clearance logger.info(f"Injecting renal clearance penalty: {penalty:.2f} for eGFR: {egfr}") - + return penalty - def hepatic_safety_adjustment(self, patient_state: Any, predicted_properties: Dict[str, float]) -> float: + def hepatic_safety_adjustment(self, patient_state: Any, predicted_properties: dict[str, float]) -> float: """ Adjusts reward for hepatic clearance if the liver is impaired. """ - ast = getattr(patient_state, 'ast', 25.0) - alt = getattr(patient_state, 'alt', 25.0) - - hepatic_toxicity = predicted_properties.get('hepatic_toxicity', 0.1) - + ast = getattr(patient_state, "ast", 25.0) + alt = getattr(patient_state, "alt", 25.0) + + hepatic_toxicity = predicted_properties.get("hepatic_toxicity", 0.1) + if ast > self.hepatic_threshold or alt > self.hepatic_threshold: # Liver is already under stress; any hepatic toxicity is amplified return -50.0 * hepatic_toxicity - + return 0.0 - def compute_n1_optimized_reward(self, - base_reward: float, - patient_state: Any, - predicted_properties: Dict[str, float]) -> float: + def compute_n1_optimized_reward( + self, base_reward: float, patient_state: Any, predicted_properties: dict[str, float] + ) -> float: """ Combines the base drug-likeness reward with patient-specific physiological penalties. """ clearance_penalty = self.dynamic_clearance_penalty(patient_state, predicted_properties) hepatic_adjustment = self.hepatic_safety_adjustment(patient_state, predicted_properties) - + # Total optimized reward total_reward = base_reward + clearance_penalty + hepatic_adjustment - + return float(total_reward) diff --git a/unicorn_platform_orchestrator.py b/unicorn_platform_orchestrator.py index 0b2c6b596a..82f9f6fe3f 100644 --- a/unicorn_platform_orchestrator.py +++ b/unicorn_platform_orchestrator.py @@ -8,6 +8,7 @@ app = typer.Typer(help="Unicorn Platform Orchestrator for ZANE") + class UnicornOrchestrator: def __init__(self): self.zkp = ZKPMarketplace() @@ -47,10 +48,11 @@ def orchestrate(self, target_seq: str): typer.echo("Unicorn workflow completed successfully!") + @app.command() def run( target_seq: str = "ATCGATCGATCG...", - fpga_device: str | None = typer.Option(None, "--fpga-device", help="FPGA/ASIC device path (sets env)") + fpga_device: str | None = typer.Option(None, "--fpga-device", help="FPGA/ASIC device path (sets env)"), ): """ Orchestrate Unicorn Modules 14-16: @@ -67,8 +69,10 @@ def run( typer.echo(f"Error: {e}", err=True) raise typer.Exit(code=1) + def main(): app() + if __name__ == "__main__": main() diff --git a/validation/breakthrough_metrics.py b/validation/breakthrough_metrics.py index 36e34123e3..7505538453 100644 --- a/validation/breakthrough_metrics.py +++ b/validation/breakthrough_metrics.py @@ -2,7 +2,7 @@ import json from pathlib import Path -from typing import Any, Dict +from typing import Any BREAKTHROUGH_BENCHMARKS = { "af3_rmsd": 1.8, # <2Å target @@ -12,7 +12,8 @@ "bbb_penetration": 0.70, } -def validate_breakthroughs(results_dir: str = "outputs/validation/2024") -> Dict[str, Any]: + +def validate_breakthroughs(results_dir: str = "outputs/validation/2024") -> dict[str, Any]: """Validate 2024 breakthrough metrics against benchmarks.""" report = {} path = Path(results_dir) @@ -23,4 +24,4 @@ def validate_breakthroughs(results_dir: str = "outputs/validation/2024") -> Dict achieved = data.get("mean", 0) passed = achieved <= threshold if "rmsd" in benchmark else achieved >= threshold report[benchmark] = {"achieved": achieved, "threshold": threshold, "passed": passed} - return report \ No newline at end of file + return report diff --git a/validation/run_enterprise_validation.py b/validation/run_enterprise_validation.py index ba04bad8de..e2637cdae0 100644 --- a/validation/run_enterprise_validation.py +++ b/validation/run_enterprise_validation.py @@ -308,7 +308,7 @@ def run_zkp_penetration_validation(rng: np.random.Generator, assets_dir: Path) - plt.ylabel("Blocked attack fraction") plt.title("ZKP Penetration Testing Attack-Block Rate") plt.grid(axis="y", alpha=0.2) - for bar, val in zip(bars, attack_values): + for bar, val in zip(bars, attack_values, strict=False): plt.text(bar.get_x() + bar.get_width() / 2, val + 0.015, f"{val:.3f}", ha="center", fontsize=8) plt.tight_layout() plt.savefig(assets_dir / "test_results_zkp_penetration.png", dpi=220) diff --git a/zane_apex_entrypoint.py b/zane_apex_entrypoint.py index 4c6f915860..19d53ded06 100644 --- a/zane_apex_entrypoint.py +++ b/zane_apex_entrypoint.py @@ -1,14 +1,16 @@ -import asyncio import argparse +import asyncio import logging -from typing import Dict, Any +from typing import Any + +from clinical.digital_twin.lab_report_parser import LabReportIngestor # Core ZANE Imports from infrastructure.knowledge_retrieval.dynamic_rag_context import DynamicTargetContext -from clinical.digital_twin.lab_report_parser import LabReportIngestor from training.de_novo_enforcer import DeNovoStrictEnforcer from training.global_reward_orchestrator import PanArchitectureReward + # Mock imports for generative models and cloud-lab API # In a full build, these would be the actual GNN/Diffusion engines async def generate_molecules_task(target_constraints, reward_orchestrator): @@ -16,17 +18,20 @@ async def generate_molecules_task(target_constraints, reward_orchestrator): logger.info("Starting de novo generation loop...") await asyncio.sleep(2) # Return a high-scoring novel molecule - return "C1=CC=C(C=C1)CC(=O)NC2=CC=CC=C2" # Mock SMILES + return "C1=CC=C(C=C1)CC(=O)NC2=CC=CC=C2" # Mock SMILES + async def dispatch_to_cloud_lab(smiles: str): """Sends the molecule blueprint to the automated synthesis API.""" logger.info(f"Dispatching molecule {smiles} to Cloud-Lab for autonomous synthesis.") return {"status": "dispatched", "job_id": "ZEN-99"} -logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') + +logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(name)s - %(levelname)s - %(message)s") logger = logging.getLogger("ZANE-APEX") -async def execute_zane_pipeline(patient_profile_path: str, target_disease: str, metadata: Dict[str, Any] = None): + +async def execute_zane_pipeline(patient_profile_path: str, target_disease: str, metadata: dict[str, Any] | None = None): """ The master chronological execution flow for the Apex N=1 pipeline. """ @@ -52,31 +57,31 @@ async def execute_zane_pipeline(patient_profile_path: str, target_disease: str, # 4. Initialize Pan-Architecture Reward Orchestrator reward_orchestrator = PanArchitectureReward(patient_state) - reward_orchestrator.novelty_enforcer = enforcer # Sync enforcer + reward_orchestrator.novelty_enforcer = enforcer # Sync enforcer # 5. Execute Generative Active Learning Loop final_smiles = await generate_molecules_task(target_constraints, reward_orchestrator) - + # 6. Final Validation & Excellence Scoring excellence_score = await reward_orchestrator.calculate_total_reward( - final_smiles, - {"docking_score": 0.95, "admet_score": 0.88} + final_smiles, {"docking_score": 0.95, "admet_score": 0.88} ) - if excellence_score > 0.8: # Threshold for "Excellence" + if excellence_score > 0.8: # Threshold for "Excellence" logger.info(f"ACHIEVED EXCELLENCE: {final_smiles} (Score: {excellence_score:.4f})") - + # 7. Dispatch to Cloud-Lab for synthesis dispatch_result = await dispatch_to_cloud_lab(final_smiles) logger.info(f"Pipeline Complete. Result: {dispatch_result}") else: logger.error("Failed to generate a molecule achieving excellence threshold.") + if __name__ == "__main__": parser = argparse.ArgumentParser(description="ZANE Apex N=1 Entrypoint") parser.add_argument("--patient_profile", type=str, required=True, help="Path to patient health report (PDF)") parser.add_argument("--target", type=str, required=True, help="Target disease/protein name") - + args = parser.parse_args() - + asyncio.run(execute_zane_pipeline(args.patient_profile, args.target))