Skip to content

engine

tit.opt.ex.engine

Single-class exhaustive search engine.

Replaces: runner.py (5 classes), roi_utils.py (4 functions). Direct SimNIBS usage — ROI metrics are computed inline.

ExSearchEngine

ExSearchEngine(leadfield_hdf: str, roi_file: str | tuple[str, int] | list[str | tuple[str, int]], roi_name: str, logger: Logger)

Exhaustive TI electrode search engine.

Owns the full pipeline: leadfield loading, ROI resolution, simulation loop, and ROI CRUD.

Source code in tit/opt/ex/engine.py
def __init__(
    self,
    leadfield_hdf: str,
    roi_file: str | tuple[str, int] | list[str | tuple[str, int]],
    roi_name: str,
    logger: logging.Logger,
):
    self.leadfield_hdf = leadfield_hdf
    self.roi_file = roi_file
    self.roi_name = roi_name
    self.logger = logger

    self.leadfield = None
    self.mesh = None
    self.idx_lf = None
    self.roi_coords = None
    self.roi_centers = None
    self.roi_indices = None
    self.roi_volumes = None
    self.gm_indices = None
    self.gm_volumes = None
    self._eval_subset = None

initialize

initialize(roi_radius: float = 3.0) -> None

Load leadfield, resolve the ROI (CSV/mask/atlas), find ROI + GM elements.

Source code in tit/opt/ex/engine.py
def initialize(self, roi_radius: float = 3.0) -> None:
    """Load leadfield, resolve the ROI (CSV/mask/atlas), find ROI + GM elements."""
    self._load_leadfield()
    self._load_roi_coordinates()
    self._find_roi_elements(roi_radius)
    self._find_gm_elements()
    self._eval_subset = None

compute_ti_fields

compute_ti_fields(e1_plus: str, e1_minus: str, e2_plus: str, e2_minus: str, current_ratios: list[tuple[float, float]]) -> list[dict[str, float]]

ROI metrics of one electrode montage at every current split.

The field is linear in the injected current, so each channel's unit-current field is gathered once and scaled per split -- the same current * (lf[a] - lf[b]) arithmetic TI.get_field performs, restricted to the elements the metrics read.

Source code in tit/opt/ex/engine.py
def compute_ti_fields(
    self,
    e1_plus: str,
    e1_minus: str,
    e2_plus: str,
    e2_minus: str,
    current_ratios: list[tuple[float, float]],
) -> list[dict[str, float]]:
    """ROI metrics of one electrode montage at every current split.

    The field is linear in the injected current, so each channel's
    unit-current field is gathered once and scaled per split -- the
    same ``current * (lf[a] - lf[b])`` arithmetic ``TI.get_field``
    performs, restricted to the elements the metrics read.
    """
    lf = self.leadfield
    idx = self.idx_lf
    subset, roi_pos, gm_pos = self._evaluation_subset()

    unit1 = TI.get_field([e1_plus, e1_minus, 1.0], lf, idx)[subset]
    unit2 = TI.get_field([e2_plus, e2_minus, 1.0], lf, idx)[subset]

    results = []
    for current_ch1_mA, current_ch2_mA in current_ratios:
        ef1 = (current_ch1_mA / 1000) * unit1
        ef2 = (current_ch2_mA / 1000) * unit2
        ti_max = TI.get_maxTI(ef1, ef2)
        data = self._roi_metrics(ti_max[roi_pos], ti_max[gm_pos])
        data["current_ch1_mA"] = current_ch1_mA
        data["current_ch2_mA"] = current_ch2_mA
        results.append(data)
    return results

compute_ti_field

compute_ti_field(e1_plus: str, e1_minus: str, current_ch1_mA: float, e2_plus: str, e2_minus: str, current_ch2_mA: float) -> dict[str, float]

Compute TI field for one montage and return ROI metrics.

Source code in tit/opt/ex/engine.py
def compute_ti_field(
    self,
    e1_plus: str,
    e1_minus: str,
    current_ch1_mA: float,
    e2_plus: str,
    e2_minus: str,
    current_ch2_mA: float,
) -> dict[str, float]:
    """Compute TI field for one montage and return ROI metrics."""
    return self.compute_ti_fields(
        e1_plus, e1_minus, e2_plus, e2_minus, [(current_ch1_mA, current_ch2_mA)]
    )[0]

run

run(e1_plus: list[str], e1_minus: list[str], e2_plus: list[str], e2_minus: list[str], current_ratios: list[tuple[float, float]], all_combinations: bool, output_dir: str, n_jobs: int = 1, symmetry_mirror_map: dict[str, str] | None = None, symmetry_pairing: str = 'within_pairs') -> dict[str, dict[str, float]]

Run the full simulation loop. Returns {mesh_key: metrics}.

Candidates are evaluated in enumeration order, one electrode montage (all its current splits) per task, on n_jobs forked workers (n_jobs < 1: the job's CPU budget, see :func:tit.opt.ex.parallel.resolve_n_jobs; 1: in-process). A symmetry_mirror_map (bucket mode) restricts the enumeration to left/right mirrored montages (see :mod:tit.opt.ex.symmetry).

Source code in tit/opt/ex/engine.py
def run(
    self,
    e1_plus: list[str],
    e1_minus: list[str],
    e2_plus: list[str],
    e2_minus: list[str],
    current_ratios: list[tuple[float, float]],
    all_combinations: bool,
    output_dir: str,
    n_jobs: int = 1,
    symmetry_mirror_map: dict[str, str] | None = None,
    symmetry_pairing: str = "within_pairs",
) -> dict[str, dict[str, float]]:
    """Run the full simulation loop. Returns {mesh_key: metrics}.

    Candidates are evaluated in enumeration order, one electrode
    montage (all its current splits) per task, on ``n_jobs`` forked
    workers (``n_jobs < 1``: the job's CPU budget, see
    :func:`tit.opt.ex.parallel.resolve_n_jobs`; ``1``: in-process).
    A *symmetry_mirror_map* (bucket mode) restricts the enumeration to
    left/right mirrored montages (see :mod:`tit.opt.ex.symmetry`).
    """
    stop = False

    def _on_signal(sig, frame):
        nonlocal stop
        stop = True

    signal.signal(signal.SIGINT, _on_signal)
    signal.signal(signal.SIGTERM, _on_signal)

    total = count_combinations(
        e1_plus,
        e1_minus,
        e2_plus,
        e2_minus,
        current_ratios,
        all_combinations,
        symmetry_mirror_map,
        symmetry_pairing,
    )
    n_jobs = resolve_n_jobs(n_jobs)
    self._log_config_summary(
        e1_plus,
        e1_minus,
        e2_plus,
        e2_minus,
        current_ratios,
        all_combinations,
        total,
        n_jobs,
        symmetry_pairing if symmetry_mirror_map is not None else None,
    )

    results: dict[str, dict[str, float]] = {}
    start_time = time.time()
    ratios = list(current_ratios)

    montages = list(
        _electrode_combinations(
            e1_plus,
            e1_minus,
            e2_plus,
            e2_minus,
            all_combinations,
            symmetry_mirror_map,
            symmetry_pairing,
        )
    )
    evaluations = evaluate_ordered(
        self,
        "compute_ti_fields",
        ((*montage, ratios) for montage in montages),
        n_jobs,
        n_tasks=len(montages),
    )

    i = 0
    for (ep1, em1, ep2, em2), montage_results in zip(montages, evaluations):
        montage_start = time.time()
        for (ch1, ch2), data in zip(ratios, montage_results):
            i += 1
            name = f"{ep1}_{em1}_and_{ep2}_{em2}_I1-{ch1:.1f}mA_I2-{ch2:.1f}mA"
            key = f"TI_field_{name}.msh"

            elapsed = time.time() - start_time
            rate = i / elapsed if elapsed > 0 else 0
            eta = (total - i) / rate if rate > 0 else 0

            self.logger.info(f"[{i}/{total}] {name}")
            self.logger.info(
                f"  {100 * i / total:.1f}% | {rate:.2f}/s | ETA {eta / 60:.1f}min"
            )

            data["electrodes"] = (ep1, em1, ep2, em2)
            results[key] = data
            self.logger.info(
                f"  {(time.time() - montage_start) / len(ratios):.2f}s | "
                f"TImax={data[f'{self.roi_name}_TImax_ROI']:.4f} "
                f"TImean={data[f'{self.roi_name}_TImean_ROI']:.4f} "
                f"Foc={data[f'{self.roi_name}_Focality']:.4f}"
            )
        if stop:
            self.logger.warning("Interrupted")
            evaluations.close()
            break

    if results:
        t = time.time() - start_time
        self.logger.info(f"\n{'=' * 60}")
        self.logger.info(
            f"Done: {len(results)}/{total} in {t / 60:.1f}min "
            f"({t / len(results):.2f}s each)"
        )
        self.logger.info(f"Output: {output_dir}")

    return results

get_available_rois staticmethod

get_available_rois(subject_id: str) -> list[str]

List ROI CSV files for a subject.

Source code in tit/opt/ex/engine.py
@staticmethod
def get_available_rois(subject_id: str) -> list[str]:
    """List ROI CSV files for a subject."""
    from tit.paths import get_path_manager

    roi_dir = Path(get_path_manager().rois(subject_id))
    return sorted(p.name for p in roi_dir.glob("*.csv"))

create_roi staticmethod

create_roi(subject_id: str, roi_name: str, x: float, y: float, z: float) -> tuple[bool, str]

Create an ROI CSV from coordinates.

Source code in tit/opt/ex/engine.py
@staticmethod
def create_roi(
    subject_id: str,
    roi_name: str,
    x: float,
    y: float,
    z: float,
) -> tuple[bool, str]:
    """Create an ROI CSV from coordinates."""
    from tit.paths import get_path_manager

    roi_dir = Path(get_path_manager().rois(subject_id))
    roi_dir.mkdir(parents=True, exist_ok=True)

    if not roi_name.endswith(".csv"):
        roi_name += ".csv"

    roi_file = roi_dir / roi_name
    with open(roi_file, "w", newline="") as f:
        csv.writer(f).writerow([x, y, z])

    roi_list = roi_dir / "roi_list.txt"
    existing = []
    if roi_list.exists():
        existing = [
            ln.strip() for ln in roi_list.read_text().splitlines() if ln.strip()
        ]
    if roi_name not in existing:
        with open(roi_list, "a") as f:
            f.write(f"{roi_name}\n")

    return True, f"ROI '{roi_name}' created at ({x:.2f}, {y:.2f}, {z:.2f})"

delete_roi staticmethod

delete_roi(subject_id: str, roi_name: str) -> tuple[bool, str]

Delete an ROI file and remove from roi_list.txt.

Source code in tit/opt/ex/engine.py
@staticmethod
def delete_roi(subject_id: str, roi_name: str) -> tuple[bool, str]:
    """Delete an ROI file and remove from roi_list.txt."""
    from tit.paths import get_path_manager

    roi_dir = Path(get_path_manager().rois(subject_id))

    if not roi_name.endswith(".csv"):
        roi_name += ".csv"

    roi_file = roi_dir / roi_name
    if roi_file.exists():
        roi_file.unlink()

    roi_list = roi_dir / "roi_list.txt"
    if roi_list.exists():
        lines = [
            ln.strip() for ln in roi_list.read_text().splitlines() if ln.strip()
        ]
        if roi_name in lines:
            lines.remove(roi_name)
            roi_list.write_text(("\n".join(lines) + "\n") if lines else "")

    return True, f"ROI '{roi_name}' deleted"

get_roi_coordinates staticmethod

get_roi_coordinates(subject_id: str, roi_name: str) -> tuple[float, float, float] | None

Read ROI center coordinates from CSV.

Source code in tit/opt/ex/engine.py
@staticmethod
def get_roi_coordinates(
    subject_id: str,
    roi_name: str,
) -> tuple[float, float, float] | None:
    """Read ROI center coordinates from CSV."""
    from tit.paths import get_path_manager

    roi_dir = Path(get_path_manager().rois(subject_id))

    if not roi_name.endswith(".csv"):
        roi_name += ".csv"

    roi_file = roi_dir / roi_name
    if not roi_file.exists():
        return None

    with open(roi_file) as f:
        for row in csv.reader(f):
            if not row:
                continue
            coords = [float(v.strip()) for v in row if v.strip()]
            if len(coords) >= 3:
                return (coords[0], coords[1], coords[2])
    return None