PhoenixAgent / app /core /reshape_executor.py
VarunRS5457
Reshape flow: redirect to classify page with AI pre-filled classifications
16d2e95
Raw
History Blame Contribute Delete
20 kB
"""Reshape-plan executor — the deterministic half of the arbitrary-layout flow.
Takes a confirmed :class:`ReshapePlan` (proposed by the agent, edited by the
user) plus the source files, and materializes the Stars output:
* a DATA matrix (rows = variables, columns = samples), and
* a METADATA table (one row per sample, with Genotype/Sex/Diet derived from
the canonical Sample_ID).
The agent decides *what* each sheet means; this module does the actual cell
arithmetic, so the transform is reproducible and auditable regardless of what
the model returned. It is intentionally tolerant: missing/again-derivable
coordinates are re-inferred so small plan errors self-heal.
Checkpoint 1 scope: samples-in-columns sheets (the ITT layout). Samples-in-rows
sheets are skipped with a warning (handled in Checkpoint 2).
"""
import logging
import re
from pathlib import Path
import pandas as pd
from app.core.layout_extractor import read_full_grid
from app.core.sas_extractor import _parse_excel_key
from app.models.mapping import ReshapePlan, SheetReshapePlan
logger = logging.getLogger(__name__)
SAMPLE_ID_NAME = "Sample_ID"
SUBJECT_TITLE_NAME = "Subject_title"
def _to_number(text: str):
"""Parse a grid cell into a float, or None if it isn't numeric."""
if text is None or text == "":
return None
try:
return float(text)
except (ValueError, TypeError):
return None
def _extract_number_token(text: str) -> str:
"""Pull the sample number out of a header cell ('858' -> '858', '858.0' -> '858')."""
text = (text or "").strip()
m = re.search(r"(\d+)(?:\.0+)?$", text)
return m.group(1) if m else text
def _norm_genotype(g: str | None) -> str:
return (g or "").strip().upper()
def _short(value: str | None, mapping: dict[str, str]) -> str:
"""Map a full form (e.g. 'Male') to its short token ('M'), case-insensitively."""
if not value:
return ""
for full, sht in mapping.items():
if value.strip().lower() == full.lower():
return sht
return value.strip()
def _timepoint_var_name(label: str) -> str | None:
"""Turn a timepoint cell ('0', '15', '30.0') into a variable name ('T0', 'T15')."""
num = _to_number(label)
if num is None:
return None
n = int(num) if float(num).is_integer() else num
return f"T{n}"
def _sample_columns(header: list[str], label_col: int) -> list[int]:
"""Contiguous non-empty cells right of the label column = the sample columns.
Stops at the first blank cell, which separates the data block from any side
block (e.g. the '% difference' table in the ITT sheets).
"""
cols: list[int] = []
for c in range(label_col + 1, len(header)):
if header[c].strip() == "":
break
cols.append(c)
return cols
def _value_rows(grid: list[list[str]], label_col: int, header_row: int,
first: int | None, last: int | None) -> list[int]:
"""Rows holding values: those below the header with a non-empty variable label."""
if first is not None and last is not None and last >= first:
return [r for r in range(first, min(last + 1, len(grid)))]
rows = []
for r in range(header_row + 1, len(grid)):
if label_col < len(grid[r]) and grid[r][label_col].strip() != "":
rows.append(r)
return rows
def _longest_ascending_run(grid: list[list[str]], col: int) -> list[int]:
"""Longest run of consecutive rows whose value in `col` strictly ascends."""
best: list[int] = []
run: list[int] = []
prev: float | None = None
for r in range(len(grid)):
v = _to_number(grid[r][col]) if col < len(grid[r]) else None
if v is not None and (prev is None or v > prev):
run.append(r)
prev = v
else:
if len(run) > len(best):
best = run
run = [r] if v is not None else []
prev = v
return best if len(best) >= len(run) else run
def _detect_time_column(grid: list[list[str]]) -> tuple[int, int, list[int]] | None:
"""Locate the TIME column of a timepoint sheet from the data itself.
The raw glucose timepoints (0,15,30,45,60,90) form the longest strictly
ascending run of any column, which pins the variable-label column, the header
row (one above the run), and the value rows without trusting model indices.
Returns (label_col, sample_id_row, value_rows) or None.
"""
best_col, best_run = None, []
width = max((len(r) for r in grid), default=0)
for c in range(width):
run = _longest_ascending_run(grid, c)
if len(run) > len(best_run):
best_col, best_run = c, run
if best_col is not None and len(best_run) >= 4:
return best_col, best_run[0] - 1, best_run
return None
def _process_columns_sheet(
grid: list[list[str]],
sheet: SheetReshapePlan,
plan: ReshapePlan,
) -> tuple[dict[str, dict[str, float]], dict[str, dict[str, str]]]:
"""Extract one samples-in-columns sheet.
Returns:
values: {sample_id -> {variable_name -> value}}
metadata: {sample_id -> {Genotype, Sex, Diet}}
"""
if not grid:
return {}, {}
# For timepoint sheets, auto-detect the TIME column and its rows straight from
# the data (robust to model coordinate errors); fall back to model indices.
detected = _detect_time_column(grid) if sheet.variable_kind == "timepoint" else None
if detected:
label_col, header_row, value_rows = detected
header_row = max(header_row, 0)
header = grid[header_row]
sample_cols = _sample_columns(header, label_col)
else:
header_row = sheet.sample_id_row if sheet.sample_id_row is not None else 0
if header_row >= len(grid):
logger.warning("%s: sample_id_row %d out of range", sheet.logical_name, header_row)
return {}, {}
header = grid[header_row]
label_col = sheet.variable_label_index
# Default the label column to the cell just left of the first sample, if unset.
if label_col is None:
first_nonempty = next((c for c, v in enumerate(header) if v.strip() != ""), 0)
label_col = max(first_nonempty, 0)
sample_cols = _sample_columns(header, label_col)
value_rows = _value_rows(grid, label_col, header_row, sheet.value_first, sheet.value_last)
if not sample_cols or not value_rows:
logger.warning("%s: no sample columns or value rows found", sheet.logical_name)
return {}, {}
geno = _norm_genotype(sheet.group.genotype)
sex_short = _short(sheet.group.sex, plan.sex_map)
diet_short = _short(sheet.group.diet, plan.diet_map)
values: dict[str, dict[str, float]] = {}
metadata: dict[str, dict[str, str]] = {}
for c in sample_cols:
raw_id = header[c]
n = _extract_number_token(raw_id)
sample_id = plan.id_format.format(
genotype=geno, sex=sex_short, diet=diet_short, n=n,
)
values.setdefault(sample_id, {})
metadata[sample_id] = {
"Genotype": geno,
"Sex": (sheet.group.sex or "").strip(),
"Diet": (sheet.group.diet or "").strip(),
}
# Raw values, keyed by variable name
raw_by_var: dict[str, float] = {}
basal: float | None = None
for r in value_rows:
label = grid[r][label_col] if label_col < len(grid[r]) else ""
val = _to_number(grid[r][c]) if c < len(grid[r]) else None
if sheet.variable_kind == "timepoint":
var = _timepoint_var_name(label)
if var is None:
continue
if basal is None:
basal = val # first timepoint is the basal reference
else:
var = label.strip()
if not var:
continue
if val is not None:
raw_by_var[var] = val
values[sample_id].update(raw_by_var)
# Computed %Basal rows for timepoint sheets
if sheet.variable_kind == "timepoint" and plan.compute_percent_basal and basal:
for r in value_rows:
label = grid[r][label_col] if label_col < len(grid[r]) else ""
tnum = _to_number(label)
if tnum is None or tnum == 0: # skip the basal timepoint itself
continue
var = _timepoint_var_name(label)
val = raw_by_var.get(var)
if val is not None:
n_label = int(tnum) if float(tnum).is_integer() else tnum
values[sample_id][f"%Basal T{n_label}"] = round(val / basal * 100, 2)
return values, metadata
# ---------------------------------------------------------------------------
# Samples-in-rows path (e.g. the insulin file: each row is a sample like
# "KO F Chow 1", columns are measurements). Side-by-side Chow/HFD blocks are
# handled as separate plan entries, each anchored on its own sample_id_col.
# ---------------------------------------------------------------------------
# Raw measurement-header aliases (normalized) -> canonical variable name.
# The same quantity is spelled many ways across sheets; all must collapse to the
# exact target-MasterSheet variable strings.
VARIABLE_ALIASES = {
"pg/ml": "pg/ml",
"80fold": "80-Fold", "80 fold": "80-Fold", "80x": "80-Fold", "80-fold": "80-Fold",
"ng/ml": "ng/ml",
"ug/ml": "ug/ml", "ug / ml": "ug/ml", "ug insulin": "ug/ml",
"dna": "DNA (ug)", "dna (ug)": "DNA (ug)", "ug dna": "DNA (ug)",
"ug dna in islets": "DNA (ug)", "dna quantity": "DNA (ug)",
"ng insulin/ug dna": "ng Ins/ug DNA", "ng ins/ug dna": "ng Ins/ug DNA",
"ng/ml/ug dna": "ng Ins/ug DNA",
"ug insulin/ug dna": "ug Ins/ug DNA", "ug insulin/ ug dna": "ug Ins/ug DNA",
"ug ins/ug dna": "ug Ins/ug DNA", "ug/ml/ug dna": "ug Ins/ug DNA",
}
_GENOTYPE_TOKENS = {"WT", "KO", "HT", "HET"}
def _norm_key(text: str) -> str:
return re.sub(r"\s+", " ", (text or "").strip().lower())
def _canonical_variable(raw: str) -> str | None:
"""Map a raw measurement-header cell to its canonical variable name.
Unknown non-empty headers pass through (trimmed) rather than being dropped,
so a novel column still becomes a variable; empty cells return None.
"""
key = _norm_key(raw)
if not key:
return None
return VARIABLE_ALIASES.get(key, raw.strip())
def _parse_row_label_id(label: str, plan: "ReshapePlan", group) -> tuple[str, str, str, str] | None:
"""Parse a sample row label like "KO F Chow 1" into (sample_id, geno, sex, diet).
Order-tolerant: tokens are classified by content, not position. Missing
tokens fall back to the sheet's `group`. Returns None if it can't produce a
genotype, sex, diet and trailing number (so callers can use it to detect
where the real data rows stop).
"""
text = (label or "").strip()
if not text:
return None
geno = sex = diet = None
n = None
for tok in re.split(r"\s+", text):
tl = tok.lower()
tu = tok.upper()
if tu in _GENOTYPE_TOKENS:
geno = "HT" if tu == "HET" else tu
elif tl in ("m", "male"):
sex = "Male"
elif tl in ("f", "female"):
sex = "Female"
elif tl == "chow":
diet = "Chow"
elif tl in ("hfd", "hf"):
diet = "HFD"
elif re.fullmatch(r"\d+", tok):
n = tok
geno = geno or (_norm_genotype(group.genotype) or None)
sex = sex or (group.sex or None)
diet = diet or (group.diet or None)
if n is None or not geno or not sex or not diet:
return None
sample_id = plan.id_format.format(
genotype=geno, sex=_short(sex, plan.sex_map), diet=_short(diet, plan.diet_map), n=n,
)
return sample_id, geno, sex, diet
def _count_alias_cells(row: list[str], label_col: int) -> int:
"""How many cells to the right of label_col are recognized measurement headers."""
return sum(
1 for c in range(label_col + 1, len(row))
if _norm_key(row[c]) in VARIABLE_ALIASES
)
def _resolve_header_row(grid, label_col: int, hint: int | None) -> int | None:
"""Find the row that actually holds measurement headers for a rows-axis block.
LLMs occasionally give a header-row index off by one; we trust the data, not
the index. Pick the row with the most alias-matching cells; fall back to the
model's hint if no row clearly qualifies.
"""
best_row, best_n = None, 1 # require >=2 recognized headers to accept
for r in range(len(grid)):
n = _count_alias_cells(grid[r], label_col)
if n > best_n:
best_row, best_n = r, n
return best_row if best_row is not None else hint
def _autodetect_label_col(grid, header_row: int, plan, group) -> int | None:
"""Pick the column whose cells below the header parse as the most sample labels."""
if header_row + 1 >= len(grid):
return None
width = max((len(r) for r in grid), default=0)
best_col, best_count = None, 0
for c in range(width):
count = sum(
1 for r in range(header_row + 1, len(grid))
if c < len(grid[r]) and _parse_row_label_id(grid[r][c], plan, group) is not None
)
if count > best_count:
best_col, best_count = c, count
return best_col
def _data_rows_for_rows_sheet(grid, label_col, header_row, first, last, plan, group) -> list[int]:
"""Rows whose label_col parses as a sample; stops at the first blank/junk row.
Honors explicit first/last when given. Otherwise scans downward from the
header, skipping leading blanks and stopping once a run of data rows ends
(which excludes the embedded summary sub-tables below the block).
"""
if first is not None and last is not None and last >= first:
return list(range(first, min(last + 1, len(grid))))
rows: list[int] = []
for r in range(header_row + 1, len(grid)):
lab = grid[r][label_col] if label_col < len(grid[r]) else ""
if _parse_row_label_id(lab, plan, group) is None:
if rows:
break
continue
rows.append(r)
return rows
def _process_rows_sheet(grid, sheet: SheetReshapePlan, plan: ReshapePlan):
"""Extract one samples-in-rows block.
Returns ({sample_id -> {variable -> value}}, {sample_id -> {Genotype,Sex,Diet}}).
"""
if not grid:
return {}, {}
hint_row = sheet.variable_label_index if sheet.variable_label_index is not None else 0
label_col = sheet.sample_id_col
if label_col is None:
label_col = _autodetect_label_col(grid, hint_row, plan, sheet.group)
if label_col is None:
logger.warning("%s: could not locate sample-label column", sheet.logical_name)
return {}, {}
# Snap to the row that actually holds measurement headers (robust to an
# off-by-one header index from the model).
header_row = _resolve_header_row(grid, label_col, hint_row)
if header_row is None or header_row >= len(grid):
logger.warning("%s: could not locate measurement header row", sheet.logical_name)
return {}, {}
header = grid[header_row]
value_cols = _sample_columns(header, label_col) # measurement columns of this block
data_rows = _data_rows_for_rows_sheet(
grid, label_col, header_row, sheet.value_first, sheet.value_last, plan, sheet.group,
)
if not value_cols or not data_rows:
logger.warning("%s: no measurement columns or sample rows found", sheet.logical_name)
return {}, {}
values: dict[str, dict[str, float]] = {}
metadata: dict[str, dict[str, str]] = {}
for r in data_rows:
lab = grid[r][label_col] if label_col < len(grid[r]) else ""
parsed = _parse_row_label_id(lab, plan, sheet.group)
if parsed is None:
continue
sample_id, geno, sex, diet = parsed
values.setdefault(sample_id, {})
metadata[sample_id] = {"Genotype": geno, "Sex": sex, "Diet": diet}
for c in value_cols:
var = _canonical_variable(header[c]) if c < len(header) else None
if not var:
continue
val = _to_number(grid[r][c]) if c < len(grid[r]) else None
if val is not None:
values[sample_id][var] = val
return values, metadata
def build_stars_from_plan(
plan: ReshapePlan,
source_paths: dict[str, Path],
) -> tuple[pd.DataFrame, pd.DataFrame]:
"""Apply a reshape plan to the source files and build the Stars data/metadata pair.
Args:
plan: the confirmed reshape plan.
source_paths: logical_name -> Path (the same map stored at upload; Excel
paths may carry a ``::SheetName`` suffix).
Returns (data_df, metadata_df) ready to be written as Stars CSVs.
"""
all_values: dict[str, dict[str, float]] = {}
all_metadata: dict[str, dict[str, str]] = {}
variable_order: list[str] = []
def _track_var(var: str) -> None:
if var not in variable_order:
variable_order.append(var)
for sheet in plan.sheets:
if not sheet.include:
continue
path = source_paths.get(sheet.logical_name)
if path is None:
logger.warning("No source path for sheet '%s'", sheet.logical_name)
continue
physical, sheet_name = _parse_excel_key(path)
grid = read_full_grid(physical, sheet_name)
if sheet.samples_axis == "rows":
values, metadata = _process_rows_sheet(grid, sheet, plan)
else:
values, metadata = _process_columns_sheet(grid, sheet, plan)
for sid, vars_ in values.items():
all_values.setdefault(sid, {}).update(vars_)
for v in vars_:
_track_var(v)
all_metadata.update(metadata)
if not all_values:
raise ValueError("Reshape produced no samples — check the plan coordinates.")
# Order variables: raw timepoints first (T0, T15…), then %Basal, then the rest.
def _var_sort_key(v: str):
m = re.fullmatch(r"T(\d+(?:\.\d+)?)", v)
if m:
return (0, float(m.group(1)))
m = re.fullmatch(r"%Basal T(\d+(?:\.\d+)?)", v)
if m:
return (1, float(m.group(1)))
return (2, variable_order.index(v))
variables = sorted(variable_order, key=_var_sort_key)
samples = list(all_values.keys())
# --- Build DATA matrix (rows = variables, columns = samples) ---
data = {sid: {var: all_values[sid].get(var) for var in variables} for sid in samples}
data_df = pd.DataFrame(data, index=variables, columns=samples)
# Prepend the Subject_title row (subject == sample here) and label the index.
subject_row = pd.DataFrame([{sid: sid for sid in samples}], index=[SUBJECT_TITLE_NAME])
data_df = pd.concat([subject_row, data_df])
data_df.index.name = SAMPLE_ID_NAME
# --- Build METADATA table ---
meta_rows = []
for sid in samples:
md = all_metadata.get(sid, {})
meta_rows.append({
SAMPLE_ID_NAME: sid,
SUBJECT_TITLE_NAME: sid,
"Genotype": md.get("Genotype", ""),
"Sex": md.get("Sex", ""),
"Diet": md.get("Diet", ""),
})
metadata_df = pd.DataFrame(meta_rows)
return data_df, metadata_df
def write_stars_csv(
data_df: pd.DataFrame,
metadata_df: pd.DataFrame,
output_dir: Path,
filename_base: str,
) -> tuple[Path, Path]:
"""Write the Stars data/metadata pair as CSV files."""
data_path = output_dir / f"data_{filename_base}.csv"
metadata_path = output_dir / f"metadata_{filename_base}.csv"
data_df.to_csv(data_path, index=True)
metadata_df.to_csv(metadata_path, index=False)
return data_path, metadata_path