Add eta interpolation; code optimization

Co-authored-by: Copilot <copilot@github.com>
This commit is contained in:
2026-06-11 16:21:47 +02:00
co-authored by Copilot
parent fd76272157
commit cc40610dc2
5 changed files with 220 additions and 72 deletions
+85 -31
View File
@@ -3,6 +3,8 @@ sys.path.append('./src')
import matplotlib
matplotlib.use('Agg')
from etaInterpolationFromPoints import interpolate_eta_from_points
from pathlib import Path
from omegaconf import OmegaConf
import torch
@@ -46,11 +48,7 @@ def apply_inverse_transforms(predictions: torch.Tensor, numberOfAugOps: int) ->
return corrected.mean(dim=1)
def prepare_output_folder(conf):
if conf.data.normalize:
normalize_suffix = '_normalized'
else:
normalize_suffix = ''
output_base = Path(conf.experiment.output_base) / conf.experiment.name / conf.model.experiment_name / f'augX{conf.inference.num_aug_ops}{normalize_suffix}'
output_base = Path(conf.experiment.output_base) / conf.experiment.name / (f'1ph_{conf.model.experiment_name}') / (f'augX{conf.inference.num_aug_ops}')
output_base.mkdir(parents=True, exist_ok=True)
OmegaConf.save(conf, output_base / 'config.yaml')
return output_base
@@ -77,41 +75,49 @@ def run_inference(model, data_loader, conf):
offset = [inputs.shape[-2] / 2., inputs.shape[-1] / 2.]
offset = torch.tensor(offset).unsqueeze(0) # (1, 2)
all_predictions += offset
print(f'mean x = {torch.mean(all_predictions[:, 0]):.4f}, std x = {torch.std(all_predictions[:, 0]):.4f}')
print(f'mean y = {torch.mean(all_predictions[:, 1]):.4f}, std y = {torch.std(all_predictions[:, 1]):.4f}')
print(f'[Inference]: mean x = {torch.mean(all_predictions[:, 0]):.4f}, std x = {torch.std(all_predictions[:, 0]):.4f}')
print(f'[Inference]: mean y = {torch.mean(all_predictions[:, 1]):.4f}, std y = {torch.std(all_predictions[:, 1]):.4f}')
return all_predictions.numpy(), all_reference_points
def accumulate_hits(predictions: np.ndarray, reference_points: np.ndarray,
binning_factor: int):
ml_super_frame = np.zeros((NY * binning_factor, NX * binning_factor))
binning_factor: int, number_of_subframes: int = 20):
ml_super_frames = np.zeros((number_of_subframes, NY * binning_factor, NX * binning_factor))
count_frame = np.zeros((NY, NX))
subpixel_dist = np.zeros((binning_factor, binning_factor))
# 绝对坐标 = 预测亚像素 + 参考点(像素左下角)
# absolute coordinate = predicted subpixel + reference point
absolute_positions = predictions + reference_points[:, :2]
# 超分辨帧 (binning)
# super resolution frames (binning)
hit_x = np.floor(absolute_positions[:, 0] * binning_factor).astype(int)
hit_y = np.floor(absolute_positions[:, 1] * binning_factor).astype(int)
np.add.at(ml_super_frame, (hit_y, hit_x), 1)
# 计数帧 (按参考点像素索引)
ref_x = (reference_points[:, 0] + 1).astype(int) # 参考点是左下角,+1 得像素索引
for i in range(number_of_subframes):
start_idx = i * len(predictions) // number_of_subframes
end_idx = min((i + 1) * len(predictions) // number_of_subframes, len(predictions))
np.add.at(ml_super_frames[i], (hit_y[start_idx:end_idx], hit_x[start_idx:end_idx]), 1)
# count frame (by reference point pixel index)
ref_x = (reference_points[:, 0] + 1).astype(int) # reference point is lower-left corner, +1 to get pixel index
ref_y = (reference_points[:, 1] + 1).astype(int)
np.add.at(count_frame, (ref_y, ref_x), 1)
# 亚像素分布
# subpixel distribution
sub_x = np.floor((absolute_positions[:, 0] % 1) * binning_factor).astype(int)
sub_y = np.floor((absolute_positions[:, 1] % 1) * binning_factor).astype(int)
np.add.at(subpixel_dist, (sub_y, sub_x), 1)
return ml_super_frame, count_frame, subpixel_dist
return ml_super_frames, count_frame, subpixel_dist
def save_results(ml_super_frame, count_frame, subpixel_dist,
def save_results(ml_super_frames, ml_super_frames_eta,
count_frame, subpixel_dist, subpixel_dist_eta,
roi: list, binning_factor: int, output_dir: Path):
x_min, x_max, y_min, y_max = roi
# 1. super-resolution frame
np.save(output_dir / '1Photon_ML_superFrames.npy', ml_super_frames)
ml_super_frame = np.sum(ml_super_frames, axis=0)
plt.figure(figsize=(8, 8))
plt.imshow(ml_super_frame[y_min*binning_factor:y_max*binning_factor, x_min*binning_factor:x_max*binning_factor], origin='lower', extent=[x_min, x_max, y_min, y_max])
plt.colorbar(label='Counts')
@@ -121,8 +127,22 @@ def save_results(ml_super_frame, count_frame, subpixel_dist,
plt.savefig(output_dir / '1Photon_ML_superFrame.png', dpi=300, bbox_inches='tight')
plt.close()
np.save(output_dir / '1Photon_ML_superFrame.npy', ml_super_frame)
# 2. count frame
# 2. super-resolution frame with eta interpolation
np.save(output_dir / '1Photon_ML_superFrames_etaInterpolated.npy', ml_super_frames_eta)
ml_super_frame_eta = np.sum(ml_super_frames_eta, axis=0)
plt.figure(figsize=(8, 8))
plt.imshow(ml_super_frame_eta[y_min*binning_factor:y_max*binning_factor, x_min*binning_factor:x_max*binning_factor], origin='lower', extent=[x_min, x_max, y_min, y_max])
plt.colorbar(label='Counts')
plt.title('ML Super-Resolution Frame with Eta Interpolation')
plt.xlabel('X (pixel)')
plt.ylabel('Y (pixel)')
plt.savefig(output_dir / '1Photon_ML_superFrame_etaInterpolated.png', dpi=300, bbox_inches='tight')
plt.close()
np.save(output_dir / '1Photon_ML_superFrame_etaInterpolated.npy', ml_super_frame_eta)
# 3. count frame
plt.figure(figsize=(8, 8))
plt.imshow(count_frame[y_min:y_max, x_min:x_max], origin='lower', extent=[x_min, x_max, y_min, y_max])
plt.colorbar(label='Counts')
@@ -133,7 +153,7 @@ def save_results(ml_super_frame, count_frame, subpixel_dist,
plt.close()
np.save(output_dir / '1Photon_count_Frame.npy', count_frame)
# 3. sub-pixel distribution
# 4. sub-pixel distribution
plt.figure(figsize=(8, 8))
plt.imshow(subpixel_dist, origin='lower',
extent=[0, binning_factor, 0, binning_factor],
@@ -146,9 +166,24 @@ def save_results(ml_super_frame, count_frame, subpixel_dist,
plt.close()
np.save(output_dir / '1Photon_subpixel_Distribution.npy', subpixel_dist)
rms, mean = np.std(subpixel_dist), np.mean(subpixel_dist)
print(f"[Plotting]: Sub-pixel distribution: RMS/Mean: {rms/mean:.4f}, expected value = {1/np.sqrt(mean):.4f} for uniform distribution")
print(f"[Plotting]: Sub-pixel distribution: RMS/Mean: {rms/mean:.4f}, expected value = {1/np.sqrt(mean):.4f} for uniform distribution")
# 5. sub-pixel distribution with eta interpolation
plt.figure(figsize=(8, 8))
plt.imshow(subpixel_dist_eta, origin='lower',
extent=[0, binning_factor, 0, binning_factor],
cmap='viridis')
plt.colorbar(label='Counts')
plt.title('Sub-pixel Distribution with Eta Interpolation')
plt.xlabel(f'Sub-pixel X (1/{binning_factor} pixel)')
plt.ylabel(f'Sub-pixel Y (1/{binning_factor} pixel)')
plt.savefig(output_dir / '1Photon_subpixel_Distribution_etaInterpolated.png', dpi=300, bbox_inches='tight')
plt.close()
np.save(output_dir / '1Photon_subpixel_Distribution_etaInterpolated.npy', subpixel_dist_eta)
rms_eta, mean_eta = np.std(subpixel_dist_eta), np.mean(subpixel_dist_eta)
print(f"[Plotting]: Sub-pixel distribution with eta interpolation: RMS/Mean: {rms_eta/mean_eta:.4f}, expected value = {1/np.sqrt(mean_eta):.4f} for uniform distribution")
print(f"Results saved to: {output_dir}")
print(f"[Plotting]: Results saved to: {output_dir}")
if __name__ == "__main__":
### output preparation
@@ -168,6 +203,9 @@ if __name__ == "__main__":
flag_normalize = conf.data.normalize
nChunks = int(np.ceil(len(files_list) / conf.inference.chunk_size))
list_of_predictions = []
list_of_reference_points = []
ml_super_frame = np.zeros((NY * BinningFactor, NX * BinningFactor))
count_frame = np.zeros((NY, NX))
subpixel_dist = np.zeros((BinningFactor, BinningFactor))
@@ -195,14 +233,30 @@ if __name__ == "__main__":
pin_memory=True
)
predictions, ref_points = run_inference(model, data_loader, conf)
_predictions, _ref_points = run_inference(model, data_loader, conf)
list_of_predictions.append(_predictions)
list_of_reference_points.append(_ref_points)
chunk_super, chunk_count, chunk_subpixel = accumulate_hits(
predictions, ref_points, binning_factor=BinningFactor
)
ml_super_frame += chunk_super
count_frame += chunk_count
subpixel_dist += chunk_subpixel
save_results(ml_super_frame, count_frame, subpixel_dist,
del dataset, data_loader
torch.cuda.empty_cache()
predictions = np.concatenate(list_of_predictions, axis=0)
ref_points = np.concatenate(list_of_reference_points, axis=0)
del list_of_predictions, list_of_reference_points
ml_super_frames, count_frame, subpixel_dist = accumulate_hits(
predictions, ref_points, binning_factor=BinningFactor
)
print('[Main]: Applying eta interpolation to predictions...')
predictions_eta_interpolated = interpolate_eta_from_points(predictions)
ml_super_frames_eta, count_frame_eta, subpixel_dist_eta = accumulate_hits(
predictions_eta_interpolated, ref_points, binning_factor=BinningFactor
)
save_results(ml_super_frames, ml_super_frames_eta,
count_frame, subpixel_dist, subpixel_dist_eta,
roi=roi, binning_factor=BinningFactor,
output_dir=output_dir)
output_dir=output_dir
)