Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ default_language_version:
python: python3.10
repos:
- repo: https://github.com/pre-commit/pre-commit-hooks
rev: v5.0.0
rev: v6.0.0
hooks:
- id: check-added-large-files
- id: check-toml
Expand All @@ -14,7 +14,7 @@ repos:
- id: end-of-file-fixer
- id: trailing-whitespace
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.8.2
rev: v0.16.4
hooks:
- id: ruff
args:
Expand All @@ -23,17 +23,17 @@ repos:
files: ^source/
- id: ruff-format
- repo: https://github.com/gitleaks/gitleaks
rev: v8.21.2
rev: v8.30.0
hooks:
- id: gitleaks
- repo: https://github.com/codespell-project/codespell
rev: v2.3.0
rev: v2.4.3
hooks:
- id: codespell
additional_dependencies:
- tomli
- repo: https://github.com/compilerla/conventional-pre-commit
rev: v3.6.0
rev: v4.4.0
hooks:
- id: conventional-pre-commit
stages: [commit-msg]
Expand Down
23 changes: 10 additions & 13 deletions source/analysis/artifacts.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
from pathlib import Path
from typing import Optional

import matplotlib.pyplot as plt
import numpy as np
Expand Down Expand Up @@ -51,7 +50,7 @@ def _set_save_path(self, dir_name: str) -> Path:
save_path.mkdir(parents=True, exist_ok=True)
return save_path

def _get_save_path(self, dir_name: str, file_name: str) -> Optional[Path]:
def _get_save_path(self, dir_name: str, file_name: str) -> Path | None:
save_path = self._artifacts_path / dir_name / file_name
if save_path.exists():
return save_path
Expand All @@ -61,25 +60,25 @@ def set_dir_params(self, params: DirParams):
self._dir_params = params
self._artifacts_path = OUT_DIR / params_to_path(params)

def save_study(self, study: Study, band_reduction: Optional[int] = None):
def save_study(self, study: Study, band_reduction: int | None = None):
self.save_metric(study.best_value, band_reduction)
self.save_params(study.best_params, band_reduction)

def save_metric(self, metric: float, band_reduction: Optional[int] = None):
def save_metric(self, metric: float, band_reduction: int | None = None):
save_path = self._set_save_path(STUDY)
if not band_reduction:
write_txt(str(metric), save_path / STUDY_BEST_METRIC)
logger.info(f"Metric saved: {metric}")

def save_params(self, params: dict, band_reduction: Optional[int] = None):
def save_params(self, params: dict, band_reduction: int | None = None):
save_path = self._set_save_path(STUDY)
if band_reduction:
write_json(params, save_path / STUDY_BEST_PARAMS_REDUCED)
else:
write_json(params, save_path / STUDY_BEST_PARAMS)
logger.info(f"Params saved: {params}")

def load_params(self, band_reduction: Optional[int] = None) -> Optional[dict]:
def load_params(self, band_reduction: int | None = None) -> dict | None:
if band_reduction:
save_path = self._get_save_path(STUDY, STUDY_BEST_PARAMS_REDUCED)
else:
Expand All @@ -100,16 +99,14 @@ def load_encoder(self) -> LabelEncoder:
"Encoder could not be found. Make sure you train the model first (cmd: train_model)"
)

def load_unfit_model(self, band_reduction: Optional[int] = None) -> BaseEstimator:
def load_unfit_model(self, band_reduction: int | None = None) -> BaseEstimator:
params = self.load_params(band_reduction)
model = import_model(self._dir_params.estimator_name)
if params:
model.set_params(**params)
return model

def save_metrics(
self, metrics: list[Metrics], band_reduction: Optional[int] = None
):
def save_metrics(self, metrics: list[Metrics], band_reduction: int | None = None):
table, metrics_all = present.generate_metrics_table(metrics)
save_path = self._set_save_path(RESULTS)
if band_reduction:
Expand All @@ -119,7 +116,7 @@ def save_metrics(
for m in metrics_all:
write_txt(f"{m.mean:.2f}", save_path / f"{m.name}")

def load_metrics(self) -> Optional[dict[str, list[Metrics]]]:
def load_metrics(self) -> dict[str, list[Metrics]] | None:
if not OUT_DIR:
logger.warning("No output directory.")
return None
Expand Down Expand Up @@ -163,7 +160,7 @@ def save_shap_values(self, values: np.ndarray):
save_path = self._set_save_path(RESULTS)
np.save(save_path / RESULT_SHAP_VALUES, values)

def load_shap_values(self) -> Optional[np.ndarray]:
def load_shap_values(self) -> np.ndarray | None:
save_path = self._get_save_path(RESULTS, RESULT_SHAP_VALUES)
if save_path:
return np.load(save_path)
Expand All @@ -173,7 +170,7 @@ def load_shap_values(self) -> Optional[np.ndarray]:
)
return None

def load_spectral_bands(self) -> Optional[np.ndarray]:
def load_spectral_bands(self) -> np.ndarray | None:
SPECTRAL_BANDS_DIR.mkdir(parents=True, exist_ok=True)
save_path = SPECTRAL_BANDS_DIR / "bands.npy"
if save_path.exists():
Expand Down
8 changes: 3 additions & 5 deletions source/analysis/extensions.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
from typing import Optional

from source.analysis.base import BaseDataLoader, BaseOptParameters, BaseSklearnModel
from source.core import settings
from source.utils.utils import import_class
Expand All @@ -13,13 +11,13 @@
DEFAULT_PARAMETERS_NAME = "ParamsSVC"


def import_data_loader(data_loader_name: Optional[str]) -> BaseDataLoader:
def import_data_loader(data_loader_name: str | None) -> BaseDataLoader:
return import_class(data_loader_name, _DATA_LOADERS_DIR, DEFAULT_DATA_LOADER_NAME)


def import_model(model_name: Optional[str]) -> BaseSklearnModel:
def import_model(model_name: str | None) -> BaseSklearnModel:
return import_class(model_name, _MODELS_DIR, DEFAULT_ESTIMATOR_NAME)


def import_parameters(parameters_name: Optional[str]) -> BaseOptParameters:
def import_parameters(parameters_name: str | None) -> BaseOptParameters:
return import_class(parameters_name, _PARAMETERS_DIR, DEFAULT_PARAMETERS_NAME)
2 changes: 1 addition & 1 deletion source/analysis/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@
@dataclass
class MetricsContainer:
def __init__(self):
self._metrics = {key: [] for key in METRIC_FUNC.keys()}
self._metrics = {key: [] for key in METRIC_FUNC}

def calculate(self, y_test: ArrayLike, y_pred: ArrayLike):
for metric_name, metric_func in METRIC_FUNC.items():
Expand Down
5 changes: 2 additions & 3 deletions source/analysis/params.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,14 @@
from pathlib import Path
from typing import Optional

from pydantic import BaseModel, model_validator
from source.analysis.extensions import DEFAULT_DATA_LOADER_NAME, DEFAULT_ESTIMATOR_NAME
from typing_extensions import Self


class DirParams(BaseModel):
estimator_name: Optional[str] = None
estimator_name: str | None = None
estimator_is_optimized: bool = False
data_loader_name: Optional[str] = None
data_loader_name: str | None = None

@model_validator(mode="after")
def set_default_names(self) -> Self:
Expand Down
12 changes: 5 additions & 7 deletions source/analysis/plots.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,7 @@
from typing import Optional

import matplotlib.cm as cm
import matplotlib.pyplot as plt
import numpy as np
import umap
from matplotlib import cm
from matplotlib.figure import Figure
from matplotlib.lines import Line2D
from sklearn.base import BaseEstimator, clone
Expand All @@ -16,7 +14,7 @@


def relevant_features(
relevances: np.ndarray, bands: Optional[np.ndarray] = None
relevances: np.ndarray, bands: np.ndarray | None = None
) -> Figure:
y = smooth_relevances(relevances)
indices_by_relevance = np.argsort(y)[::-1]
Expand Down Expand Up @@ -55,7 +53,7 @@ def relevant_features(


def relevant_amplitudes(
relevances: np.ndarray, bands: Optional[np.ndarray] = None
relevances: np.ndarray, bands: np.ndarray | None = None
) -> Figure:
if bands is None:
bands = np.arange(len(relevances))
Expand Down Expand Up @@ -121,7 +119,7 @@ def signatures_display(
encoder: LabelEncoder,
X: np.ndarray,
y: np.ndarray,
bands: Optional[np.ndarray] = None,
bands: np.ndarray | None = None,
*,
x_label: str = "Spectral bands",
y_label: str = "",
Expand Down Expand Up @@ -186,7 +184,7 @@ def umap_display(
encoder: LabelEncoder,
X: np.ndarray,
y: np.ndarray,
meta: Optional[np.ndarray] = None,
meta: np.ndarray | None = None,
) -> Figure:
y_encoded = np.array(encoder.fit_transform(y))
classes = encoder.classes_
Expand Down
7 changes: 3 additions & 4 deletions source/analysis/present.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
from collections import defaultdict
from typing import Optional

from rich import print as RichPrint
from rich.text import Text as RichText
Expand All @@ -22,9 +21,9 @@ def generate_metrics_table(metrics: list[Metrics]):

def display_metrics(
metrics: dict[str, list[Metrics]],
model: Optional[str] = None,
do_optimize: Optional[bool] = None,
data_loader: Optional[str] = None,
model: str | None = None,
do_optimize: bool | None = None,
data_loader: str | None = None,
):
def check_filter(params: DirParams) -> bool:
if model is not None and params.estimator_name != model:
Expand Down
25 changes: 12 additions & 13 deletions source/cli/analysis.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
from functools import reduce
from typing import Optional

import typer
from rich import print
Expand All @@ -19,7 +18,7 @@

@app.command()
def test_load_data(
data_loader: Optional[str] = None,
data_loader: str | None = None,
):
loader = import_data_loader(data_loader)
loader.load_data()
Expand All @@ -32,10 +31,10 @@ def test_load_data(

@app.command()
def train_model(
model: Optional[str] = None,
data_loader: Optional[str] = None,
model: str | None = None,
data_loader: str | None = None,
do_optimize: bool = False,
parameters: Optional[list[str]] = None,
parameters: list[str] | None = None,
):
artifacts.set_dir_params(
DirParams(
Expand Down Expand Up @@ -69,8 +68,8 @@ def train_model(

@app.command()
def generate_metrics(
model: Optional[str] = None,
data_loader: Optional[str] = None,
model: str | None = None,
data_loader: str | None = None,
do_optimize: bool = False,
):
artifacts.set_dir_params(
Expand All @@ -93,8 +92,8 @@ def generate_metrics(

@app.command()
def generate_plots(
model: Optional[str] = None,
data_loader: Optional[str] = None,
model: str | None = None,
data_loader: str | None = None,
do_optimize: bool = False,
):
artifacts.set_dir_params(
Expand Down Expand Up @@ -127,8 +126,8 @@ def generate_plots(

@app.command()
def calculate_relevances(
model: Optional[str] = None,
data_loader: Optional[str] = None,
model: str | None = None,
data_loader: str | None = None,
do_optimize: bool = False,
):
artifacts.set_dir_params(
Expand All @@ -149,8 +148,8 @@ def calculate_relevances(

@app.command()
def display_metrics(
model: Optional[str] = None,
data_loader: Optional[str] = None,
model: str | None = None,
data_loader: str | None = None,
do_optimize: bool = False,
):
metrics_ = artifacts.load_metrics()
Expand Down
3 changes: 1 addition & 2 deletions source/cli/misc.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
from source.misc.display_image import (
display_spectral_image,
)
from typing import Optional

app = typer.Typer()

Expand All @@ -14,5 +13,5 @@ def check_images():


@app.command()
def display_image(label: Optional[str] = None):
def display_image(label: str | None = None):
display_spectral_image(label)
20 changes: 9 additions & 11 deletions source/cli/segment.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,4 @@
from typing import Optional

import typer
from source.segmentation.artifacts import (
load_all_selected_areas,
load_model,
load_transformation_matrix,
save_model,
save_selected_areas,
save_transformation_matrix,
)
from source.misc.display_image import (
display_spectral_images_with_areas,
)
Expand All @@ -21,6 +11,14 @@
select_areas_on_images,
train_xgboost_model,
)
from source.segmentation.artifacts import (
load_all_selected_areas,
load_model,
load_transformation_matrix,
save_model,
save_selected_areas,
save_transformation_matrix,
)

app = typer.Typer()

Expand Down Expand Up @@ -52,7 +50,7 @@ def train_model():


@app.command()
def segment_images(label: Optional[str] = None):
def segment_images(label: str | None = None):
encoder_cam1, model_cam1, encoder_cam2, model_cam2 = load_model()
matx = load_transformation_matrix()
perform_segmentation(
Expand Down
2 changes: 1 addition & 1 deletion source/core/logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ def emit(self, record: logging.LogRecord) -> None:
logger_opt.log(record.levelname, record.getMessage())


@lru_cache()
@lru_cache
def setup_logger(debug: bool = False) -> loguru._Logger: # type: ignore
LOGGING_LEVEL = logging.DEBUG if debug else logging.INFO
logging.basicConfig(
Expand Down
4 changes: 2 additions & 2 deletions source/misc/check_images.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,6 @@ def check_spectral_images():
f" Number of labels: {len(labels_cam1)} (cam1), {len(labels_cam2)} (cam2)\n"
)
msg += f" Number of unique labels: {len(labels_unique)} \n"
msg += f" Labels: \n{str(labels_unique)} \n"
msg += f" Duplicated labels: \n{str(labels_duplicated)} \n"
msg += f" Labels: \n{labels_unique!s} \n"
msg += f" Duplicated labels: \n{labels_duplicated!s} \n"
logger.info(f"Report: \n{msg}")
Loading
Loading