"""Training, preprocessing, and prediction helpers for FTIR model routes."""
from __future__ import annotations
import os
import re
import time
from dataclasses import replace
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterable
from .._paths import find_project_root
# Tool caches live in the user's project directory (never in the package install location,
# which may be a read-only site-packages when xpectra is installed from PyPI).
CACHE_DIR = find_project_root() / ".cache"
os.environ.setdefault("MPLCONFIGDIR", str(CACHE_DIR / "matplotlib"))
os.environ.setdefault("XDG_CACHE_HOME", str(CACHE_DIR))
os.environ.setdefault("NUMBA_CACHE_DIR", str(CACHE_DIR / "numba"))
os.environ.setdefault("NUMBA_DISABLE_JIT", "1")
for _cache_path in (
CACHE_DIR,
CACHE_DIR / "matplotlib",
CACHE_DIR / "numba",
):
_cache_path.mkdir(parents=True, exist_ok=True)
import joblib
import numpy as np
import pandas as pd
from sklearn.base import clone
from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import LabelEncoder, StandardScaler
import xpectrass
from xpectrass import FTIRdataanalysis, FTIRdataprocessing, combine_datasets
from .config import (
CORE_GROUP,
ORDERED_GROUPS,
PE_MEMBERS,
PET_LIKE,
PP_MEMBERS,
PS_LIKE,
ROUTE_FILES,
ROUTES,
PreprocessConfig,
)
[docs]
def safe_name(name: str) -> str:
"""Return a stable filesystem/column-safe version of a model name."""
return re.sub(r"[^A-Za-z0-9]+", "_", name).strip("_").lower()
[docs]
def normalize_routes(routes: Iterable[str] | None) -> list[str]:
if not routes or "all" in routes:
return list(ROUTES)
normalized = []
for route in routes:
if route not in ROUTES:
raise ValueError(f"Unknown route '{route}'. Expected one of {ROUTES} or 'all'.")
if route not in normalized:
normalized.append(route)
return normalized
[docs]
def read_csv(path: str | Path) -> pd.DataFrame:
return pd.read_csv(path, compression="infer")
[docs]
def is_wavenumber_column(column: Any) -> bool:
try:
float(str(column))
except (TypeError, ValueError):
return False
return True
[docs]
def spectral_columns_sorted(
df: pd.DataFrame,
wn_min: float | None = None,
wn_max: float | None = None,
) -> tuple[list[str], np.ndarray]:
pairs: list[tuple[float, str]] = []
for column in df.columns:
if not is_wavenumber_column(column):
continue
wn = float(str(column))
if wn_min is not None and wn < wn_min:
continue
if wn_max is not None and wn > wn_max:
continue
pairs.append((wn, str(column)))
if not pairs:
raise ValueError("No spectral wavenumber columns were found.")
pairs.sort(key=lambda item: item[0])
wavenumbers = np.array([wn for wn, _ in pairs], dtype=float)
columns = [column for _, column in pairs]
return columns, wavenumbers
[docs]
def map_polymer_type(polymer: Any) -> str:
if pd.isna(polymer):
return "Other"
polymer = str(polymer).strip()
if polymer in CORE_GROUP:
return polymer
if polymer in PE_MEMBERS:
return "PE"
if polymer in PP_MEMBERS:
return "PP"
if polymer in PS_LIKE:
return "PS"
if polymer in PET_LIKE:
return "PET"
return "Other"
[docs]
def load_training_dataframe(
route: str,
processed_dir: str | Path,
config: PreprocessConfig,
) -> pd.DataFrame:
if route not in ROUTE_FILES:
raise ValueError(f"Unknown route '{route}'. Expected one of {ROUTES}.")
path = Path(processed_dir) / ROUTE_FILES[route]
if not path.exists():
raise FileNotFoundError(f"Missing processed route data: {path}")
df = read_csv(path)
if "study" in df.columns:
# The route files carry all datasets (labeled cores, the unknown pool, and the
# external augmentation sources). Supervised training uses only the labeled cores;
# externals must be pooled explicitly per experiment, never implicitly.
df = df[df["study"].isin(config.labeled_studies)].copy()
else:
df = df.copy()
if config.label_column not in df.columns:
raise ValueError(f"Training data is missing label column '{config.label_column}'.")
df["type_original"] = df[config.label_column]
df[config.label_column] = df["type_original"].map(map_polymer_type)
df[config.label_column] = pd.Categorical(
df[config.label_column],
categories=ORDERED_GROUPS,
ordered=True,
)
return df
[docs]
def training_arrays(
df: pd.DataFrame,
config: PreprocessConfig,
) -> dict[str, Any]:
feature_columns, wavenumbers = spectral_columns_sorted(df, config.wn_min, config.wn_max)
X_raw = df[feature_columns].to_numpy(dtype=float)
X_raw = np.nan_to_num(X_raw, nan=0.0, posinf=0.0, neginf=0.0)
labels = df[config.label_column].astype(str).to_numpy()
label_encoder = LabelEncoder()
y = label_encoder.fit_transform(labels)
return {
"X_raw": X_raw,
"y": y,
"label_encoder": label_encoder,
"class_names": label_encoder.classes_.tolist(),
"feature_columns": feature_columns,
"wavenumbers": wavenumbers,
}
[docs]
def available_models(df: pd.DataFrame, config: PreprocessConfig) -> dict[str, Any]:
fda = FTIRdataanalysis(
df=df,
dataset_name="Combined dataset",
label_column=config.label_column,
sample_id_column=config.sample_id_column,
exclude_columns=["study", "sample_id", "environmental", "resolution", "type_original"],
random_state=config.random_state,
n_jobs=config.n_jobs,
)
return fda.models
[docs]
def select_models(
models: dict[str, Any],
requested: Iterable[str] | None,
limit: int | None = None,
) -> dict[str, Any]:
if not requested or "all" in requested:
selected = dict(models)
else:
selected = {}
safe_lookup = {safe_name(name): name for name in models}
for model_name in requested:
resolved = models.get(model_name)
if resolved is None:
real_name = safe_lookup.get(safe_name(model_name))
if real_name is not None:
resolved = models[real_name]
model_name = real_name
if resolved is None:
available = ", ".join(models)
raise ValueError(f"Model '{model_name}' not found. Available models: {available}")
selected[model_name] = resolved
if limit is not None:
selected = dict(list(selected.items())[:limit])
return selected
[docs]
def evaluate_holdout(
estimator: Any,
X_raw: np.ndarray,
y: np.ndarray,
config: PreprocessConfig,
) -> dict[str, Any]:
X_train, X_test, y_train, y_test = train_test_split(
X_raw,
y,
test_size=config.test_size,
random_state=config.random_state,
stratify=y,
)
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)
model = clone(estimator)
started = time.time()
model.fit(X_train_scaled, y_train)
train_time = time.time() - started
started = time.time()
y_pred = model.predict(X_test_scaled)
pred_time = time.time() - started
y_train_pred = model.predict(X_train_scaled)
return {
"status": "success",
"fit_scope": "holdout",
"train_samples": int(len(y_train)),
"test_samples": int(len(y_test)),
"train_time": train_time,
"pred_time": pred_time,
"test_accuracy": accuracy_score(y_test, y_pred),
"test_precision": precision_score(y_test, y_pred, average="weighted", zero_division=0),
"test_recall": recall_score(y_test, y_pred, average="weighted", zero_division=0),
"test_f1": f1_score(y_test, y_pred, average="weighted", zero_division=0),
"train_accuracy": accuracy_score(y_train, y_train_pred),
}
[docs]
def fit_final_model(
estimator: Any,
X_raw: np.ndarray,
y: np.ndarray,
config: PreprocessConfig,
fit_on: str = "full",
) -> tuple[Any, StandardScaler, dict[str, Any]]:
if fit_on not in {"full", "split"}:
raise ValueError("fit_on must be 'full' or 'split'.")
if fit_on == "split":
X_fit, _, y_fit, _ = train_test_split(
X_raw,
y,
test_size=config.test_size,
random_state=config.random_state,
stratify=y,
)
else:
X_fit, y_fit = X_raw, y
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X_fit)
model = clone(estimator)
started = time.time()
model.fit(X_scaled, y_fit)
train_time = time.time() - started
fit_info = {
"fit_scope": fit_on,
"fit_samples": int(len(y_fit)),
"fit_time": train_time,
}
return model, scaler, fit_info
[docs]
def artifact_path(models_dir: str | Path, route: str, model_name: str) -> Path:
return Path(models_dir) / route / f"{safe_name(model_name)}.joblib"
[docs]
def make_artifact(
route: str,
model_name: str,
model: Any,
scaler: StandardScaler,
arrays: dict[str, Any],
config: PreprocessConfig,
source_file: str | Path,
metrics: dict[str, Any] | None,
fit_info: dict[str, Any],
) -> dict[str, Any]:
return {
"format_version": 1,
"created_at": datetime.now(timezone.utc).isoformat(),
"route": route,
"model_name": model_name,
"model_safe_name": safe_name(model_name),
"model": model,
"scaler": scaler,
"label_encoder": arrays["label_encoder"],
"class_names": arrays["class_names"],
"feature_columns": arrays["feature_columns"],
"wavenumbers": arrays["wavenumbers"],
"preprocessing": config.to_dict(),
"label_grouping": {
"PE": sorted(PE_MEMBERS),
"PET": sorted(PET_LIKE),
"PP": sorted(PP_MEMBERS),
"PS": sorted(PS_LIKE),
"PVC": sorted(CORE_GROUP),
"Other": "all remaining labels",
},
"training": {
"source_file": str(source_file),
"unknown_study_excluded": config.unknown_study,
**fit_info,
},
"metrics": metrics or {},
"versions": {
"xpectrass": getattr(xpectrass, "__version__", "unknown"),
"numpy": np.__version__,
"pandas": pd.__version__,
},
}
[docs]
def train_route(
route: str,
processed_dir: str | Path | None = None,
models_dir: str | Path | None = None,
model_names: Iterable[str] | None = None,
config: PreprocessConfig | None = None,
fit_on: str = "full",
evaluate: bool = False,
skip_existing: bool = False,
limit_models: int | None = None,
) -> list[dict[str, Any]]:
config = config or PreprocessConfig()
root = find_project_root()
processed_dir = Path(processed_dir) if processed_dir else root / "processed_data"
models_dir = Path(models_dir) if models_dir else root / "models"
df = load_training_dataframe(route, processed_dir, config)
arrays = training_arrays(df, config)
models = select_models(available_models(df, config), model_names, limit=limit_models)
route_records: list[dict[str, Any]] = []
route_dir = Path(models_dir) / route
route_dir.mkdir(parents=True, exist_ok=True)
for model_name, estimator in models.items():
out_path = artifact_path(models_dir, route, model_name)
record = {
"route": route,
"model_name": model_name,
"artifact": str(out_path),
"status": "pending",
}
if skip_existing and out_path.exists():
record["status"] = "skipped_existing"
route_records.append(record)
continue
try:
metrics = evaluate_holdout(estimator, arrays["X_raw"], arrays["y"], config) if evaluate else {}
model, scaler, fit_info = fit_final_model(
estimator,
arrays["X_raw"],
arrays["y"],
config,
fit_on=fit_on,
)
source_file = Path(processed_dir) / ROUTE_FILES[route]
artifact = make_artifact(
route=route,
model_name=model_name,
model=model,
scaler=scaler,
arrays=arrays,
config=config,
source_file=source_file,
metrics=metrics,
fit_info=fit_info,
)
joblib.dump(artifact, out_path)
record.update({"status": "success", **fit_info})
if metrics:
record.update(
{
"test_accuracy": metrics.get("test_accuracy"),
"test_f1": metrics.get("test_f1"),
}
)
except Exception as exc: # Keep long all-model runs moving.
record.update({"status": "failed", "error": str(exc)})
route_records.append(record)
return route_records
[docs]
def train_routes(
routes: Iterable[str] | None = None,
processed_dir: str | Path | None = None,
models_dir: str | Path | None = None,
model_names: Iterable[str] | None = None,
config: PreprocessConfig | None = None,
fit_on: str = "full",
evaluate: bool = False,
skip_existing: bool = False,
limit_models: int | None = None,
) -> pd.DataFrame:
records: list[dict[str, Any]] = []
for route in normalize_routes(routes):
records.extend(
train_route(
route=route,
processed_dir=processed_dir,
models_dir=models_dir,
model_names=model_names,
config=config,
fit_on=fit_on,
evaluate=evaluate,
skip_existing=skip_existing,
limit_models=limit_models,
)
)
return pd.DataFrame(records)
[docs]
def preprocess_raw_dataframe(
df: pd.DataFrame,
config: PreprocessConfig | None = None,
force_absorbance: bool = False,
absorbance_scale_factor: float | None = None,
plot: bool = False,
) -> pd.DataFrame:
config = config or PreprocessConfig()
df = ensure_prediction_metadata(df, config)
if force_absorbance or absorbance_scale_factor is not None:
df_abs = force_absorbance_input(df, scale_factor=absorbance_scale_factor)
fdp = FTIRdataprocessing(
df=df_abs,
label_column=config.label_column,
sample_id_column=config.sample_id_column,
exclude_regions=list(config.exclude_regions),
interpolate_regions=list(config.interpolate_regions),
flat_windows=list(config.flat_windows),
random_state=config.random_state,
n_jobs=config.n_jobs,
)
denoised = fdp.denoise_spect(data=df_abs, method=config.denoising_method, plot=plot)
baseline = fdp.correct_baseline(data=denoised, method=config.baseline_method, plot=plot)
atm = fdp.exclude_interpolate(data=baseline, method=config.interpolate_method, plot=plot)
return fdp.normalize(data=atm, method=config.normalization_method, plot=plot)
fdp = FTIRdataprocessing(
df=df,
label_column=config.label_column,
sample_id_column=config.sample_id_column,
exclude_regions=list(config.exclude_regions),
interpolate_regions=list(config.interpolate_regions),
flat_windows=list(config.flat_windows),
random_state=config.random_state,
n_jobs=config.n_jobs,
)
return fdp._get_normalized_data(
denoising_method=config.denoising_method,
baseline_correction_method=config.baseline_method,
interpolate_method=config.interpolate_method,
normalization_method=config.normalization_method,
plot=plot,
)
[docs]
def combine_normalized_to_grid(
df: pd.DataFrame,
config: PreprocessConfig | None = None,
study_name: str = "prediction",
) -> pd.DataFrame:
config = config or PreprocessConfig()
df = ensure_prediction_metadata(df, config)
if "study" in df.columns:
df = df.rename(columns={"study": "source_study"})
combined, _ = combine_datasets(
datasets=[df],
wn_min=config.wn_min,
wn_max=config.wn_max,
resolution=config.resolution,
descending=config.descending,
method=config.combine_method,
label_column=config.label_column,
sample_id_column=config.sample_id_column,
exclude_columns=None,
add_study_column=True,
study_names=[study_name],
show_progress=False,
n_jobs=1,
grid_mode="custom",
report_coverage=False,
data_mode="normalized",
)
return combined
[docs]
def derivative_route_dataframe(
normalized_df: pd.DataFrame,
order: int,
config: PreprocessConfig | None = None,
) -> pd.DataFrame:
config = config or PreprocessConfig()
fdp = FTIRdataprocessing(
df=normalized_df,
label_column=config.label_column,
sample_id_column=config.sample_id_column,
exclude_regions=list(config.exclude_regions),
interpolate_regions=list(config.interpolate_regions),
flat_windows=list(config.flat_windows),
random_state=config.random_state,
n_jobs=config.n_jobs,
)
return fdp.derivatives(
data=normalized_df,
order=order,
window_length=config.derivative_window_length,
polyorder=config.derivative_polyorder,
delta=config.derivative_delta,
plot=False,
save_plot=False,
save_path=None,
)
[docs]
def make_prediction_route_dataframes(
input_df: pd.DataFrame,
routes: Iterable[str] | None = None,
config: PreprocessConfig | None = None,
input_stage: str = "raw",
force_absorbance: bool = False,
absorbance_scale_factor: float | None = None,
save_features_dir: str | Path | None = None,
) -> dict[str, pd.DataFrame]:
config = config or PreprocessConfig()
requested_routes = normalize_routes(routes)
if input_stage not in {"raw", "normalized", "route"}:
raise ValueError("input_stage must be 'raw', 'normalized', or 'route'.")
if input_stage == "route":
route_dfs = {
route: ensure_prediction_metadata(input_df, config).copy()
for route in requested_routes
}
else:
if input_stage == "raw":
normalized = preprocess_raw_dataframe(
input_df,
config=config,
force_absorbance=force_absorbance,
absorbance_scale_factor=absorbance_scale_factor,
plot=False,
)
else:
normalized = input_df.copy()
normalized = combine_normalized_to_grid(normalized, config=config)
route_dfs = {}
if "norm" in requested_routes:
route_dfs["norm"] = normalized
if "deriv1" in requested_routes:
route_dfs["deriv1"] = derivative_route_dataframe(normalized, order=1, config=config)
if "deriv2" in requested_routes:
route_dfs["deriv2"] = derivative_route_dataframe(normalized, order=2, config=config)
if save_features_dir is not None:
features_dir = Path(save_features_dir)
features_dir.mkdir(parents=True, exist_ok=True)
for route, route_df in route_dfs.items():
route_df.to_csv(features_dir / f"{route}_features.csv", index=False)
return route_dfs
[docs]
def artifact_matches(artifact: dict[str, Any], path: Path, requested: Iterable[str] | None) -> bool:
if not requested or "all" in requested:
return True
requested_safe = {safe_name(item) for item in requested}
requested_raw = set(requested)
return (
artifact.get("model_name") in requested_raw
or artifact.get("model_safe_name") in requested_safe
or path.stem in requested_safe
)
[docs]
def load_artifacts(
models_dir: str | Path,
route: str,
model_names: Iterable[str] | None = None,
) -> list[tuple[Path, dict[str, Any]]]:
route_dir = Path(models_dir) / route
if not route_dir.exists():
return []
artifacts: list[tuple[Path, dict[str, Any]]] = []
for path in sorted(route_dir.glob("*.joblib")):
artifact = joblib.load(path)
if artifact.get("route") != route:
continue
if artifact_matches(artifact, path, model_names):
artifacts.append((path, artifact))
return artifacts
[docs]
def align_features_to_artifact(
df: pd.DataFrame,
artifact: dict[str, Any],
) -> np.ndarray:
target_wavenumbers = np.asarray(artifact["wavenumbers"], dtype=float)
spectral_cols, wavenumbers = spectral_columns_sorted(df)
lookup = {round(float(wn), 6): col for col, wn in zip(spectral_cols, wavenumbers)}
ordered_cols: list[str] = []
missing: list[float] = []
for wn in target_wavenumbers:
col = lookup.get(round(float(wn), 6))
if col is None:
missing.append(float(wn))
else:
ordered_cols.append(col)
if missing:
first_missing = ", ".join(f"{wn:.4f}" for wn in missing[:8])
raise ValueError(
f"Input data is missing {len(missing)} required wavenumber columns "
f"for route '{artifact.get('route')}'. First missing: {first_missing}"
)
X_raw = df[ordered_cols].to_numpy(dtype=float)
X_raw = np.nan_to_num(X_raw, nan=0.0, posinf=0.0, neginf=0.0)
return artifact["scaler"].transform(X_raw)
[docs]
def predict_with_artifact(
df: pd.DataFrame,
artifact: dict[str, Any],
include_probabilities: bool = False,
) -> pd.DataFrame:
X = align_features_to_artifact(df, artifact)
model = artifact["model"]
label_encoder = artifact["label_encoder"]
y_pred = model.predict(X).astype(int)
labels = label_encoder.inverse_transform(y_pred)
model_key = f"{artifact['route']}_{artifact['model_safe_name']}"
output = pd.DataFrame({f"pred_{model_key}": labels})
if hasattr(model, "predict_proba"):
try:
proba = model.predict_proba(X)
output[f"confidence_{model_key}"] = np.max(proba, axis=1)
if include_probabilities:
for idx, class_name in enumerate(label_encoder.classes_):
output[f"proba_{model_key}_{safe_name(str(class_name))}"] = proba[:, idx]
except Exception:
output[f"confidence_{model_key}"] = np.nan
else:
output[f"confidence_{model_key}"] = np.nan
return output
[docs]
def predict_csv(
input_csv: str | Path,
output_csv: str | Path,
routes: Iterable[str] | None = None,
models_dir: str | Path | None = None,
model_names: Iterable[str] | None = None,
config: PreprocessConfig | None = None,
input_stage: str = "raw",
force_absorbance: bool = False,
absorbance_scale_factor: float | None = None,
include_probabilities: bool = False,
save_features_dir: str | Path | None = None,
) -> pd.DataFrame:
config = config or PreprocessConfig()
models_dir = Path(models_dir) if models_dir else find_project_root() / "models"
input_df = read_csv(input_csv)
route_dfs = make_prediction_route_dataframes(
input_df,
routes=routes,
config=config,
input_stage=input_stage,
force_absorbance=force_absorbance,
absorbance_scale_factor=absorbance_scale_factor,
save_features_dir=save_features_dir,
)
if not route_dfs:
raise ValueError("No route dataframes were generated for prediction.")
first_df = next(iter(route_dfs.values()))
prediction_output = metadata_frame(first_df)
loaded_any = False
for route, route_df in route_dfs.items():
artifacts = load_artifacts(models_dir, route, model_names=model_names)
if not artifacts:
raise FileNotFoundError(
f"No saved artifacts found for route '{route}' in {Path(models_dir) / route}. "
"Train models first with scripts/train_models.py."
)
for _, artifact in artifacts:
loaded_any = True
preds = predict_with_artifact(
route_df,
artifact,
include_probabilities=include_probabilities,
)
prediction_output = pd.concat([prediction_output, preds], axis=1)
if not loaded_any:
raise FileNotFoundError("No matching model artifacts were loaded.")
output_csv = Path(output_csv)
output_csv.parent.mkdir(parents=True, exist_ok=True)
prediction_output.to_csv(output_csv, index=False)
return prediction_output
[docs]
def config_with_overrides(
config: PreprocessConfig | None = None,
**overrides: Any,
) -> PreprocessConfig:
config = config or PreprocessConfig()
clean = {key: value for key, value in overrides.items() if value is not None}
return replace(config, **clean)