TIDE / src /modules /pubtime /comparator.py
hiasgnpsadgd's picture
Deploy TIDE Docker backend
33d7314 verified
Raw
History Blame Contribute Delete
9.96 kB
from __future__ import annotations
from pathlib import Path
from typing import Any
import pandas as pd
from .adapter import prepare_from_archive, prepare_from_source
from .common import profile_for_pubtime
COMPARATOR_FIELDS = (
"phase",
"primary_purpose",
"intervention_model",
"allocation",
"masking",
"gender",
"number_of_arms",
"has_dmc",
"has_us_facility",
)
def summarize_from_archive(profile: dict[str, Any], project_root: Path) -> dict[str, Any]:
prepared, source_label, source_type = prepare_from_archive(profile.get("domain", ""), project_root)
return _summarize_prepared(profile, prepared, source_label, source_type)
def summarize_from_zip(profile: dict[str, Any], zip_path: Path) -> dict[str, Any]:
prepared, source_label, source_type = prepare_from_source(profile.get("domain", ""), zip_path, zip_path.name)
return _summarize_prepared(profile, prepared, source_label, source_type)
def _summarize_prepared(
profile: dict[str, Any],
prepared,
source_label: str,
source_type: str,
) -> dict[str, Any]:
rows = _frame_to_rows(prepared.survival_df)
pubtime_profile = profile_for_pubtime(profile)
matched = _matched_rows(pubtime_profile, rows)
comparator_rows = matched if len(matched) >= 20 else rows
return {
"domain": prepared.domain,
"source_archive": source_label,
"source_type": source_type,
"raw_rows": prepared.raw_rows,
"analytic_rows": prepared.analytic_rows,
"domain_rows": prepared.survival_rows,
"matched_rows": len(matched),
"used_rows": len(comparator_rows),
"match_strategy": _match_strategy(len(matched)),
"pubtime_preparation": {
"predictors": list(prepared.predictors),
"high_missing_predictors": list(prepared.high_missing_predictors),
"cox_model": prepared.cox_model_status,
},
"summary": _summarize_rows(comparator_rows),
"comparison": _compare_profile(profile, comparator_rows),
"examples": _example_rows(pubtime_profile, comparator_rows),
}
def _frame_to_rows(frame: pd.DataFrame) -> tuple[dict[str, Any], ...]:
if frame.empty:
return tuple()
wanted = {
"nct_id",
"brief_title",
"phase",
"primary_purpose",
"intervention_model",
"allocation",
"masking",
"enrollment",
"actual_duration",
"number_of_facilities",
"number_of_arms",
"number_of_primary_outcomes_to_measure",
"number_of_secondary_outcomes_to_measure",
"has_dmc",
"has_us_facility",
"were_results_reported",
"pub_date",
"event_pub",
"time_to_pub",
"result_count",
}
available = [column for column in wanted if column in frame.columns]
rows = frame[available].where(pd.notna(frame[available]), "").to_dict("records")
return tuple(rows)
def _matched_rows(profile: dict[str, Any], rows: tuple[dict[str, Any], ...]) -> list[dict[str, Any]]:
matches: list[dict[str, Any]] = []
for row in rows:
score = _match_score(profile, row)
if score >= 3:
row_with_score = dict(row)
row_with_score["_match_score"] = score
matches.append(row_with_score)
matches.sort(key=lambda row: int(row.get("_match_score", 0)), reverse=True)
return matches
def _match_score(profile: dict[str, Any], row: dict[str, Any]) -> int:
score = 0
for field in COMPARATOR_FIELDS:
profile_value = profile.get(field)
row_value = row.get(field)
if profile_value is None or profile_value == "" or row_value == "":
continue
if field == "number_of_arms":
if _to_float(profile_value) == _to_float(row_value):
score += 1
continue
if str(profile_value).upper() == str(row_value).upper():
score += 1
return score
def _match_strategy(match_count: int) -> str:
if match_count >= 20:
return "matched_pubtime_prepared_fields"
return "domain_fallback_too_few_matches"
def _summarize_rows(rows: list[dict[str, Any]] | tuple[dict[str, Any], ...]) -> dict[str, Any]:
enrollments = [_to_float(row.get("enrollment")) for row in rows]
durations = [_to_float(row.get("actual_duration")) for row in rows]
facilities = [_to_float(row.get("number_of_facilities")) for row in rows]
arms = [_to_float(row.get("number_of_arms")) for row in rows]
primary_outcomes = [_to_float(row.get("number_of_primary_outcomes_to_measure")) for row in rows]
secondary_outcomes = [_to_float(row.get("number_of_secondary_outcomes_to_measure")) for row in rows]
time_to_pub = [_to_float(row.get("time_to_pub")) for row in rows]
return {
"median_enrollment": _median(enrollments),
"median_duration_months": _median(durations),
"median_facilities": _median(facilities),
"median_arms": _median(arms),
"median_primary_outcomes": _median(primary_outcomes),
"median_secondary_outcomes": _median(secondary_outcomes),
"median_time_to_publication_days": _median(time_to_pub),
"publication_rate": _rate(rows, _has_publication),
"results_reported_rate": _rate(rows, lambda row: _is_true(row.get("were_results_reported"))),
"dmc_rate": _rate(rows, lambda row: _is_true(row.get("has_dmc"))),
"us_facility_rate": _rate(rows, lambda row: _is_true(row.get("has_us_facility"))),
}
def _compare_profile(
profile: dict[str, Any], rows: list[dict[str, Any]] | tuple[dict[str, Any], ...]
) -> dict[str, Any]:
summary = _summarize_rows(rows)
flags: list[str] = []
enrollment = _to_float(profile.get("enrollment"))
median_enrollment = summary["median_enrollment"]
if enrollment and median_enrollment:
if enrollment >= median_enrollment * 1.5:
flags.append("Planned enrollment is substantially above the comparator median.")
elif enrollment <= median_enrollment * 0.5:
flags.append("Planned enrollment is substantially below the comparator median.")
facilities = _to_float(profile.get("number_of_facilities"))
median_facilities = summary["median_facilities"]
if facilities and median_facilities and facilities >= max(5, median_facilities * 2):
flags.append("Planned site count is high relative to comparators.")
arms = _to_float(profile.get("number_of_arms"))
median_arms = summary["median_arms"]
if arms and median_arms and arms >= max(4, median_arms * 2):
flags.append("Planned arm count is high relative to comparators.")
primary_outcomes = _to_float(profile.get("number_of_primary_outcomes"))
median_primary = summary["median_primary_outcomes"]
if primary_outcomes and median_primary and primary_outcomes >= max(3, median_primary * 2):
flags.append("Primary outcome count is high relative to comparators.")
secondary_outcomes = _to_float(profile.get("number_of_secondary_outcomes"))
median_secondary = summary["median_secondary_outcomes"]
if secondary_outcomes and median_secondary and secondary_outcomes >= max(6, median_secondary * 2):
flags.append("Secondary outcome count is high relative to comparators.")
dmc_rate = summary["dmc_rate"]
if profile.get("has_dmc") is False and dmc_rate is not None and dmc_rate >= 0.5:
flags.append("Most comparable trials have a DMC; this design does not.")
if profile.get("allocation") == "RANDOMIZED" and profile.get("masking") == "NONE":
flags.append("Randomized open-label design should be reviewed for bias and ascertainment risk.")
if profile.get("criteria") and len(str(profile["criteria"]).split()) < 20:
flags.append("Eligibility criteria are too short for strong complexity assessment.")
priority = "standard"
if len(flags) >= 3:
priority = "high"
elif flags:
priority = "focused"
return {
"review_priority": priority,
"flags": flags,
"note": "This is a PubTime-prepared historical comparator, not a validated prediction model.",
}
def _example_rows(
profile: dict[str, Any], rows: list[dict[str, Any]] | tuple[dict[str, Any], ...]
) -> list[dict[str, Any]]:
examples: list[dict[str, Any]] = []
ranked = sorted(rows, key=lambda row: _match_score(profile, row), reverse=True)
for row in ranked[:5]:
examples.append(
{
"nct_id": row.get("nct_id"),
"brief_title": row.get("brief_title"),
"phase": row.get("phase"),
"enrollment": _to_float(row.get("enrollment")),
"facilities": _to_float(row.get("number_of_facilities")),
"arms": _to_float(row.get("number_of_arms")),
"published": _has_publication(row),
}
)
return examples
def _median(values: list[float | None]) -> float | None:
cleaned = [value for value in values if value is not None]
if not cleaned:
return None
return round(float(pd.Series(cleaned).median()), 2)
def _rate(rows: list[dict[str, Any]] | tuple[dict[str, Any], ...], predicate) -> float | None:
if not rows:
return None
return round(sum(1 for row in rows if predicate(row)) / len(rows), 3)
def _to_float(value: Any) -> float | None:
if value is None or value == "":
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def _is_true(value: Any) -> bool:
return str(value).strip().lower() == "true"
def _has_publication(row: dict[str, Any]) -> bool:
if row.get("event_pub") != "":
return _to_float(row.get("event_pub")) == 1
if row.get("pub_date"):
return True
result_count = _to_float(row.get("result_count"))
return bool(result_count and result_count > 0)