Files
UAV-CO2/src/gasflux/processing_pipelines.py
2026-04-20 09:43:21 +08:00

322 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import logging
from datetime import datetime
from pathlib import Path
from scipy import stats
import pandas as pd
import yaml
from src.gasflux import background,plotting,processing,reporting,interpolation,pre_processing,gas
from abc import ABC, abstractmethod
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
def read_csv(file_path: Path) -> pd.DataFrame:
"""Read a CSV file and return a DataFrame."""
try:
return pd.read_csv(file_path)
except FileNotFoundError:
logger.exception(f"File not found: {file_path}")
raise
def load_config(config_path: Path) -> dict:
"""Load a YAML config file and return a dictionary."""
try:
with open(config_path) as file:
return yaml.safe_load(file)
except FileNotFoundError:
logger.exception(f"Config file not found: {config_path}")
raise
except yaml.YAMLError as e:
logger.exception(f"Error parsing YAML config: {e}")
raise
class DataValidator: # TODO(me): decide whether to move this to preprocessing
"""Validate the data before processing."""
def __init__(self, df: pd.DataFrame, config: dict) -> None:
"""Initialise the validator with the DataFrame and config."""
self.df = df
self.required_cols = config["required_cols"].copy()
self.required_cols.update(config["gases"])
def validate(self) -> None:
"""Validate the data."""
self._check_is_df()
self._check_cols()
self._check_dtypes()
self._check_ranges()
logger.info("Data validation passed")
def _check_is_df(self) -> None:
"""Check that the input is a DataFrame."""
if not isinstance(self.df, pd.DataFrame):
logging.error("Input data is not a DataFrame.")
raise ValueError("Input data is not a DataFrame.")
def _check_cols(self) -> None:
"""Check that the required columns are present."""
missing_cols = set(self.required_cols) - set(self.df.columns)
if missing_cols:
logging.error(f"Missing or mislabelled columns: {missing_cols}")
raise ValueError(f"Missing or mislabelled columns: {missing_cols}")
def _check_dtypes(self) -> None:
"""Check that the required columns are of the correct type and do not contain NaN values."""
for col in self.required_cols:
if col in self.df:
if self.df[col].isna().any():
logging.error(f"Column '{col}' contains NaN values.")
raise ValueError(f"Column '{col}' contains NaN values.")
if self.df[col].dtype != "float64":
logging.error(f"Column '{col}' is not of type 'float64'.")
raise ValueError(f"Column '{col}' is not of type 'float64'.")
def _check_ranges(self) -> None:
"""Check that the required columns are within the specified ranges."""
for col, (min_val, max_val) in self.required_cols.items():
if col in self.df.columns and not self.df[col].between(min_val, max_val, inclusive="both").all():
logging.error(f"Column '{col}' contains values out of range: {min_val} to {max_val}.")
raise ValueError(f"Column '{col}' contains values out of range: {min_val} to {max_val}.")
class BackgroundStrategy(ABC):
def __init__(self, data_processor):
self.data_processor = data_processor
@abstractmethod
def process(self):
pass
class AlgorithmicBaselineStrategy(BackgroundStrategy):
def process(self):
logger.info("Applying algorithmic background correction")
for gas in self.data_processor.gases:
(
self.data_processor.df,
self.data_processor.figs["background"][gas],
self.data_processor.text[f"background_{gas}"],
) = background.algorithmic_baseline(
df=self.data_processor.df,
gas=gas,
algorithmic_baseline_settings=self.data_processor.config["algorithmic_baseline_settings"],
)
self.data_processor.df_std = self.data_processor.df.copy()
class SensorStrategy(ABC):
def __init__(self, data_processor):
self.data_processor = data_processor
@abstractmethod
def process(self):
pass
class InSituSensorStrategy(SensorStrategy):
def process(self):
logger.info("Processing in-situ (point) data")
for gas in self.data_processor.gases:
self.data_processor.figs["scatter_3d"][gas] = None
self.data_processor.figs["windrose"] = None
self.data_processor.figs["wind_timeseries"] = None
class SpatialProcessingStrategy(ABC):
def __init__(self, data_processor):
self.data_processor = data_processor
@abstractmethod
def process(self):
pass
class CurtainSpatialProcessingStrategy(SpatialProcessingStrategy):
def process(self):
logger.info("Applying curtain spatial processing")
self.data_processor.dfs["original"] = self.data_processor.df.copy()
self.data_processor.df, self.data_processor.start_transect, self.data_processor.end_transect = (
processing.largest_monotonic_transect_series(self.data_processor.df)
)
self.data_processor.dfs["removed"] = self.data_processor.dfs["original"].loc[
self.data_processor.dfs["original"].index.difference(self.data_processor.df.index)
]
self.data_processor.df, self.data_processor.plane_angle = processing.flatten_linear_plane(
self.data_processor.df
)
self.data_processor.df = processing.wind_offset_correction(
self.data_processor.df, self.data_processor.plane_angle
)
for gas_name in self.data_processor.gases:
#计算通量
self.data_processor.df = gas.gas_flux_column(self.data_processor.df, gas_name)
fig = self.data_processor.figs["scatter_3d"].get(gas_name)
if fig is not None:
fig.add_trace(
go.Scatter3d(
x=self.data_processor.dfs["removed"]["utm_easting"],
y=self.data_processor.dfs["removed"]["utm_northing"],
z=self.data_processor.dfs["removed"]["height_ato"],
mode="markers",
marker={"size": 2, "color": "black", "symbol": "circle", "opacity": 0.5},
)
)
class SpiralSpatialProcessingStrategy(SpatialProcessingStrategy):
def process(self):
logger.info("Applying spiral spatial processing")
self.data_processor.dfs["original"] = self.data_processor.df.copy()
# self.data_processor.df, self.data_processor.start_transect, self.data_processor.end_transect = (
# gasflux.processing.largest_monotonic_transect_series(self.data_processor.df)
# no wind offset correction - assume wind is perpendicular to the spiral
self.data_processor.dfs["removed"] = self.data_processor.dfs["original"].loc[
self.data_processor.dfs["original"].index.difference(self.data_processor.df.index)
]
(
self.data_processor.df,
self.data_processor.circle_radius,
self.data_processor.circle_center_x,
self.data_processor.circle_center_y,
) = processing.circle_deviation(self.data_processor.df, x_col="utm_easting", y_col="utm_northing")
# 使用配置中第一个气体的归一化列名
primary_gas = self.data_processor.gases[0]
y_col = f"{primary_gas}_normalised"
self.data_processor.df = processing.recentre_azimuth(
self.data_processor.df, r=self.data_processor.circle_radius, y=y_col
)
self.data_processor.df["x"] = self.data_processor.df["circumference_distance"]
for gas_name in self.data_processor.gases:
self.data_processor.df = gas.gas_flux_column(self.data_processor.df, gas_name)
fig = self.data_processor.figs["scatter_3d"].get(gas_name)
if fig is not None:
fig.add_trace(
go.Scatter3d(
x=self.data_processor.dfs["removed"]["utm_easting"],
y=self.data_processor.dfs["removed"]["utm_northing"],
z=self.data_processor.dfs["removed"]["height_ato"],
mode="markers",
marker={"size": 2, "color": "black", "symbol": "circle", "opacity": 0.5},
)
)
class InterpolationStrategy(ABC):
def __init__(self, data_processor):
self.data_processor = data_processor
@abstractmethod
def process(self):
pass
class KrigingInterpolationStrategy(InterpolationStrategy):
def process(self):
logger.info("Applying kriging interpolation")
for gas in self.data_processor.gases:
(
self.data_processor.output_vars["krig_parameters"][gas],
self.data_processor.text[f"krig_output_{gas}"],
self.data_processor.figs["contour"][gas],
self.data_processor.figs["krig_grid"][gas],
self.data_processor.figs["semivariogram"][gas],
) = interpolation.ordinary_kriging(
df=self.data_processor.df,
x="x",
y="height_ato",
gas=gas,
ordinary_kriging_settings=self.data_processor.config["ordinary_kriging_settings"],
**self.data_processor.config["semivariogram_settings"],
)
logger.info(f"Kriged {gas}")
class DataProcessor:
def __init__(self, config: dict, df: pd.DataFrame):
self.config: dict = config
self.df: pd.DataFrame = df
self.gases: list[str] = list(config["gases"].keys())
self.processing_time = datetime.now()
self.figs: dict = {
"scatter_3d": {},
"windrose": None,
"wind_timeseries": None,
"background": {},
"contour": {},
"krig_grid": {},
"semivariogram": {},
}
self.text: dict = {}
self.output_vars: dict = {"krig_parameters": {}, "std": {}}
self.dfs: dict = {}
self.reports: dict = {}
def strategy_selection(self):
self.background_strategy: BackgroundStrategy
if self.config["strategies"]["background"] == "algorithm":
self.background_strategy = AlgorithmicBaselineStrategy(self)
self.sensor_strategy: SensorStrategy
if self.config["strategies"]["sensor"] == "insitu":
self.sensor_strategy = InSituSensorStrategy(self)
self.spatial_processing_strategy: SpatialProcessingStrategy
if self.config["strategies"]["spatial"] == "curtain":
self.spatial_processing_strategy = CurtainSpatialProcessingStrategy(self)
if self.config["strategies"]["spatial"] == "spiral":
self.spatial_processing_strategy = SpiralSpatialProcessingStrategy(self)
self.interpolation_strategy: InterpolationStrategy
if self.config["strategies"]["interpolation"] == "kriging":
self.interpolation_strategy = KrigingInterpolationStrategy(self)
def process(self):
self.df = pre_processing.add_utm(self.df)
self.df = pre_processing.add_course(self.df)
DataValidator(self.df, self.config).validate()
self.background_strategy.process()
self.sensor_strategy.process()
self.spatial_processing_strategy.process()
self.interpolation_strategy.process()
# Reporting
for gas in self.gases:
self.reports[gas] = reporting.mass_balance_report(
krig_params=self.output_vars["krig_parameters"][gas],
wind_fig=self.figs["wind_timeseries"],
background_fig=self.figs["background"][gas],
threed_fig=self.figs["scatter_3d"][gas],
krig_fig=self.figs["contour"][gas],
windrose_fig=self.figs["windrose"],
)
# Collecting descriptive variables
self.output_vars["std"]["windspeed"] = self.df["windspeed"].std()
self.output_vars["std"]["windddir"] = stats.circstd(self.df["winddir"], high=360)
for gas in self.gases:
self.output_vars["std"][f"{gas}_background"] = self.df.loc[
~self.df[f"{gas}_signal"], f"{gas}_normalised"
].std()
def process_main(data_file: Path, config_file: Path, output_dir: Path, task_id: str | None = None) -> DataProcessor:
"""Main function to run the pipeline."""
config = load_config(config_file)
# 优先使用 task_id,否则退回文件 stem
name = task_id if task_id else data_file.stem
df = read_csv(data_file)
processor = DataProcessor(config, df)
processor.strategy_selection()
processor.process()
reporting.generate_reports(name, processor, config, output_dir)
logger.info("Processing complete")
return processor # 返回processor对象