Remote_sensing / pca_module.py
DrAmar's picture
Upload 4 files
cdeafc5 verified
Raw History Blame Contribute Delete
7.67 kB
import os
import numpy as np
import geopandas as gpd
import pandas as pd
import rasterio
from rasterio.features import shapes
from sklearn.decomposition import PCA
import logging
from tqdm import tqdm
from shapely.geometry import shape
# Configure logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
# Choose one of the following band mappings based on selected bands for PCA
band_mapping_case1 = {2: 0, 4: 1, 5: 2, 6: 3} # For Bands 2, 4, 5, and 6
band_mapping_case2 = {2: 0, 5: 1, 6: 2, 7: 3} # For Bands 2, 5, 6, and 7
# Choose band mapping (adjust as needed)
band_mapping = band_mapping_case1
# Function to read specific band file
def read_band_file(input_directory, scene_name, band_index):
band_filename = f"{scene_name}_SR_B{band_index}.TIF"
band_path = os.path.join(input_directory, band_filename)
if not os.path.exists(band_path):
logging.error(f"File not found: {band_path}")
return None
with rasterio.open(band_path) as src:
return src.read(1), src.transform, src.crs
# Function to stack bands for PCA
def stack_bands(scene_name, input_directory, output_directory, bands):
stacked_bands, transform, crs = [], None, None
for band_index in bands:
result = read_band_file(input_directory, scene_name, band_index)
if result is not None:
band_data, transform, crs = result
stacked_bands.append(band_data)
else:
logging.warning(f"Skipping band {band_index} for scene {scene_name} due to missing file.")
if len(stacked_bands) == len(bands):
stacked_array = np.stack(stacked_bands, axis=-1)
stacked_path = os.path.join(output_directory, f"{scene_name}_stacked.tif")
with rasterio.open(stacked_path, 'w', driver='GTiff', height=stacked_array.shape[0],
width=stacked_array.shape[1], count=len(bands), dtype=stacked_array.dtype,
crs=crs, transform=transform) as dst:
for i in range(len(bands)):
dst.write(stacked_array[..., i], i + 1)
logging.info(f"Stacked bands saved at: {stacked_path}")
return stacked_path
else:
logging.warning(f"Could not stack all bands for scene: {scene_name}")
return None
# Function to perform PCA on stacked bands
def perform_pca(stacked_path):
with rasterio.open(stacked_path) as src:
data = src.read()
reshaped_data = data.reshape(data.shape[0], -1).T
pca = PCA(n_components=min(data.shape[0], data.shape[2]))
pca_results = pca.fit_transform(reshaped_data)
loadings_df = pd.DataFrame(pca.components_.T, columns=[f'PC{i + 1}' for i in range(pca.n_components_)])
return pca_results.reshape(data.shape[1], data.shape[2], -1), loadings_df
# Function to analyze principal components
def analyze_principal_components(loadings, criteria):
selected_pcs = {}
for mineral, bands in criteria.items():
found = False
for pc_index in range(loadings.shape[0]):
if all((loadings[pc_index, band_mapping[band_index]] > 0 if sign == '+' else loadings[pc_index, band_mapping[band_index]] < 0)
for band_index, sign in bands.items() if band_index in band_mapping):
selected_pcs[mineral] = pc_index + 1
logging.info(f"Selected PC for {mineral} is PC{pc_index + 1}.")
found = True
break
if not found:
logging.warning(f"No suitable PC found for {mineral}.")
return selected_pcs
# Function to log PCA statistics
def log_pca_stats(loadings_df, selected_pcs, output_directory, scene_name):
stats = []
for mineral, pc_index in selected_pcs.items():
pc_column = loadings_df[f'PC{pc_index}']
mean_val = pc_column.mean()
std_val = pc_column.std()
logging.info(f"{mineral} - PC{pc_index}: Mean={mean_val}, Std Dev={std_val}")
stats.append({
'Mineral': mineral,
'Selected_PC': f'PC{pc_index}',
'Mean_Loading': mean_val,
'Std_Dev_Loading': std_val,
**{f'Band_{col}': loadings_df.iloc[band_mapping[col], pc_index - 1] for col in band_mapping.keys()}
})
stats_df = pd.DataFrame(stats)
stats_csv_path = os.path.join(output_directory, f"{scene_name}_PCA_Loading_Stats.csv")
stats_df.to_csv(stats_csv_path, index=False)
logging.info(f"PCA loading values and statistics saved at: {stats_csv_path}")
# Function to apply thresholds and convert to vector
def apply_threshold_and_convert_to_vector(data, mean, std_dev, threshold_levels, output_shapefile_base, transform, crs):
for label, multiplier in threshold_levels.items():
threshold = mean + multiplier * std_dev
mask = data > threshold
with rasterio.Env():
shapes_gen = shapes(mask.astype(np.uint8), transform=transform)
polygons = [shape(geom) for geom, val in shapes_gen if val == 1]
shapefile_path = f"{output_shapefile_base}_{label}.shp"
gdf = gpd.GeoDataFrame(geometry=polygons, crs=crs)
gdf.to_file(shapefile_path)
logging.info(f"Shapefile created for threshold {label} at: {shapefile_path}")
# Main function to process a single scene
def process_single_scene(scene_name, input_directory, output_directory, mineral_criteria, threshold_levels):
bands = [2, 4, 5, 6] if band_mapping == band_mapping_case1 else [2, 5, 6, 7]
stacked_path = stack_bands(scene_name, input_directory, output_directory, bands)
if stacked_path is not None:
with rasterio.open(stacked_path) as src:
transform, crs = src.transform, src.crs
pca_results, loadings_df = perform_pca(stacked_path)
selected_pcs = analyze_principal_components(loadings_df.values, mineral_criteria)
log_pca_stats(loadings_df, selected_pcs, output_directory, scene_name)
for mineral, pc_index in selected_pcs.items():
pc_data = pca_results[:, :, pc_index - 1]
mean, std_dev = np.mean(pc_data), np.std(pc_data)
shapefile_path_base = os.path.join(output_directory, f"{scene_name}_{mineral}")
apply_threshold_and_convert_to_vector(pc_data, mean, std_dev, threshold_levels, shapefile_path_base, transform, crs)
# Batch process function for all files in the directory
def batch_process_scenes(input_directory, output_directory, mineral_criteria, threshold_levels):
processed_scenes = set()
tif_files = [file for file in os.listdir(input_directory) if file.endswith('.TIF')]
for file in tqdm(tif_files, desc="Processing scenes"):
scene_name = "_".join(file.split("_")[:-2])
if scene_name not in processed_scenes:
process_single_scene(scene_name, input_directory, output_directory, mineral_criteria, threshold_levels)
processed_scenes.add(scene_name)
# Main execution block for standalone use
if __name__ == "__main__":
input_directory = r'D:\RS\Arabian Shield\Complied data\Data for BR\segregated_files\SR'
output_directory = r'D:\RS\Arabian Shield\Complied data\Processed_PCA'
mineral_criteria = {
'Iron_Oxides': {2: '+', 4: '-'},
'Hydroxyl_Alteration': {2: '+', 5: '+', 6: '-', 7: '+'},
'Limonite': {2: '+', 5: '-'}
}
threshold_levels = {'98%': 3, '95%': 2, '92%': 1}
batch_process_scenes(input_directory, output_directory, mineral_criteria, threshold_levels)
logging.info("Processing completed!")