detect-cme / app.py
broadfield-dev's picture
Update app.py
2aa7ef4 verified
Raw History Blame
13.8 kB
import gradio as gr
import numpy as np
import cv2
from PIL import Image
import requests
from datetime import datetime, timedelta
import io
import os
from urllib.parse import urljoin
# Default parameters
low_int = 10
high_int = 100
edge_thresh = 50
accum_thresh = 45
center_tol = 30
morph_dia = 5
min_rad = 70
def fetch_sdo_images(start_date, end_date, ident="0171", size="1024", tool="hmiigr"):
"""Fetch SDO images from NASA URL for a given date range."""
try:
start = datetime.strptime(start_date, "%Y-%m-%d %H:%M:%S")
end = datetime.strptime(end_date, "%Y-%m-%d %H:%M:%S")
if start > end:
return None, "Start date must be before end date."
base_url = "https://sdo.gsfc.nasa.gov/assets/img/browse/"
frames = []
current = start
while current <= end:
date_str = current.strftime("%Y%m%d_%H%M%S")
year, month, day = current.strftime("%Y"), current.strftime("%m"), current.strftime("%d")
url = urljoin(base_url, f"{year}/{month}/{day}/{date_str}_{ident}_{size}_{tool}.jpg")
try:
response = requests.get(url, timeout=5)
if response.status_code == 200:
img = Image.open(io.BytesIO(response.content)).convert('L') # Convert to grayscale
frames.append(np.array(img))
else:
print(f"Failed to fetch {url}: Status {response.status_code}")
except Exception as e:
print(f"Error fetching {url}: {str(e)}")
current += timedelta(minutes=12) # SDO images are typically 12 minutes apart
if not frames:
return None, "No images found in the specified date range."
return frames, None
except Exception as e:
return None, f"Error fetching images: {str(e)}"
def extract_frames(gif_path):
"""Extract frames from a GIF and return as a list of numpy arrays."""
try:
img = Image.open(gif_path)
frames = []
while True:
frame = img.convert('L') # Convert to grayscale
frames.append(np.array(frame))
try:
img.seek(img.tell() + 1)
except EOFError:
break
return frames, None
except Exception as e:
return None, f"Error loading GIF: {str(e)}"
def preprocess_frame(frame, lower_bound, upper_bound, morph_iterations):
"""Preprocess a frame: isolate mid-to-light pixels and enhance circular patterns."""
blurred = cv2.GaussianBlur(frame, (9, 9), 0)
mask = cv2.inRange(blurred, lower_bound, upper_bound)
kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5))
enhanced = cv2.dilate(mask, kernel, iterations=morph_iterations)
return enhanced
def detect_circles(frame_diff, image_center, center_tolerance, param1, param2, min_radius=20, max_radius=200):
"""Detect circles in a frame difference image, centered at the Sun."""
circles = cv2.HoughCircles(
frame_diff,
cv2.HOUGH_GRADIENT,
dp=1.5,
minDist=100,
param1=param1,
param2=param2,
minRadius=min_radius,
maxRadius=max_radius
)
if circles is not None:
circles = np.round(circles[0, :]).astype("int")
filtered_circles = []
for (x, y, r) in circles:
if (abs(x - image_center[0]) < center_tolerance and
abs(y - image_center[1]) < center_tolerance):
filtered_circles.append((x, y, r))
return filtered_circles if filtered_circles else None
return None
def create_gif(frames, output_path, duration=0.5):
"""Create a GIF from a list of frames."""
pil_frames = [Image.fromarray(frame) for frame in frames]
pil_frames[0].save(
output_path,
save_all=True,
append_images=pil_frames[1:],
duration=int(duration * 1000), # Duration in milliseconds
loop=0
)
return output_path
def handle_fetch(start_date, end_date, ident, size, tool):
"""Fetch SDO images and return frames for preview and state."""
frames, error = fetch_sdo_images(start_date, end_date, ident, size, tool)
if error:
return error, [], frames
preview_frames = [Image.fromarray(frame) for frame in frames]
return f"Fetched {len(frames)} images successfully.", preview_frames, frames
def analyze_images(frames, lower_bound, upper_bound, param1, param2, center_tolerance, morph_iterations, min_rad, display_mode):
"""Analyze frames for concentric circles, highlighting growing series."""
try:
if not frames or len(frames) < 2:
return "At least 2 frames are required for analysis.", [], None
# Determine image center
height, width = frames[0].shape
image_center = (width // 2, height // 2)
min_radius = int(min_rad)
max_radius = min(height, width) // 2
# Process frames and detect circles
all_circle_data = []
for i in range(len(frames) - 1):
frame1 = preprocess_frame(frames[i], lower_bound, upper_bound, morph_iterations)
frame2 = preprocess_frame(frames[i + 1], lower_bound, upper_bound, morph_iterations)
frame_diff = cv2.absdiff(frame2, frame1)
frame_diff = cv2.convertScaleAbs(frame_diff, alpha=3.0, beta=0)
circles = detect_circles(frame_diff, image_center, center_tolerance, param1, param2, min_radius, max_radius)
if circles:
largest_circle = max(circles, key=lambda c: c[2])
x, y, r = largest_circle
all_circle_data.append({
"frame": i + 1,
"center": (x, y),
"radius": r,
"output_frame": frames[i + 1]
})
# Find growing series
growing_circle_data = []
current_series = []
if all_circle_data:
current_series.append(all_circle_data[0])
for i in range(1, len(all_circle_data)):
if all_circle_data[i]["radius"] > current_series[-1]["radius"]:
current_series.append(all_circle_data[i])
else:
if len(current_series) > len(growing_circle_data):
growing_circle_data = current_series.copy()
current_series = [all_circle_data[i]]
if len(current_series) > len(growing_circle_data):
growing_circle_data = current_series.copy()
growing_frames = set(c["frame"] for c in growing_circle_data)
results = []
report = f"Analysis Report (as of {datetime.now().strftime('%I:%M %p PDT, %B %d, %Y')}):\n"
# Prepare output based on display mode
if display_mode == "All Frames":
results = [Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_GRAY2RGB)) for frame in frames]
elif display_mode == "Detected Frames":
for c in all_circle_data:
output_frame = cv2.cvtColor(c["output_frame"], cv2.COLOR_GRAY2RGB)
cv2.circle(output_frame, c["center"], c["radius"], (0, 255, 0), 2) # Green for detected
if c["frame"] in growing_frames:
cv2.circle(output_frame, c["center"], c["radius"] + 2, (255, 165, 0), 2) # Orange for growing
results.append(Image.fromarray(output_frame))
elif display_mode == "Both (Detected Replaces Original)":
for i, frame in enumerate(frames):
if i + 1 in [c["frame"] for c in all_circle_data]:
for c in all_circle_data:
if c["frame"] == i + 1:
output_frame = cv2.cvtColor(c["output_frame"], cv2.COLOR_GRAY2RGB)
cv2.circle(output_frame, c["center"], c["radius"], (0, 255, 0), 2) # Green
if c["frame"] in growing_frames:
cv2.circle(output_frame, c["center"], c["radius"] + 2, (255, 165, 0), 2) # Orange
results.append(Image.fromarray(output_frame))
break
else:
results.append(Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_GRAY2RGB)))
# Generate report
if all_circle_data:
report += f"\nAll Frames with Detected Circles ({len(all_circle_data)} frames):\n"
for c in all_circle_data:
report += f"Frame {c['frame']}: Center at {c['center']}, Radius {c['radius']} pixels\n"
else:
report += "No circles detected.\n"
if growing_circle_data:
report += f"\nSeries of Frames with Growing Circles ({len(growing_circle_data)} frames):\n"
for c in growing_circle_data:
report += f"Frame {c['frame']}: Center at {c['center']}, Radius {c['radius']} pixels\n"
report += "\nConclusion: Growing concentric circles detected, indicative of a potential Earth-directed CME."
else:
report += "\nNo growing concentric circles detected. CME may not be Earth-directed."
# Create GIF if results exist
gif_path = None
if results:
gif_frames = [np.array(img) for img in results]
gif_path = "output.gif"
create_gif(gif_frames, gif_path)
return report, results, gif_path
except Exception as e:
return f"Error during analysis: {str(e)}", [], None
def process_input(gif_file, start_date, end_date, ident, size, tool, lower_bound, upper_bound, param1, param2, center_tolerance, morph_iterations, min_rad, display_mode, fetched_frames_state):
"""Process either uploaded GIF or fetched SDO images."""
if gif_file:
frames, error = extract_frames(gif_file.name)
if error:
return error, [], None, []
else:
frames = fetched_frames_state
if not frames:
return "No fetched frames available. Please fetch images first.", [], None, []
# Preview all frames
preview = [Image.fromarray(frame) for frame in frames] if frames else []
# Analyze frames
report, results, gif_path = analyze_images(
frames, lower_bound, upper_bound, param1, param2, center_tolerance, morph_iterations, min_rad, display_mode
)
return report, results, gif_path, preview
# Gradio Blocks interface
with gr.Blocks(title="Solar CME Detection") as demo:
gr.Markdown("""
# Solar CME Detection
Upload a GIF or fetch SDO images by date range to detect concentric circles indicative of coronal mass ejections (CMEs).
Green circles mark detected features; orange circles highlight growing series (potential Earth-directed CMEs).
""")
# State to store fetched frames
fetched_frames_state = gr.State(value=[])
with gr.Row():
with gr.Column():
gr.Markdown("### Input Options")
gif_input = gr.File(label="Upload Solar GIF (optional)", file_types=[".gif"])
start_date = gr.Textbox(label="Start Date (YYYY-MM-DD HH:MM:SS)", value="2025-05-24 00:00:00")
end_date = gr.Textbox(label="End Date (YYYY-MM-DD HH:MM:SS)", value="2025-05-24 23:59:59")
ident = gr.Textbox(label="Image Identifier", value="0171")
size = gr.Textbox(label="Image Size", value="1024")
tool = gr.Textbox(label="Instrument", value="hmiigr")
fetch_button = gr.Button("Fetch Images from URL")
gr.Markdown("### Analysis Parameters")
lower_bound = gr.Slider(minimum=0, maximum=255, value=low_int, step=1, label="Lower Intensity Bound (0-255)")
upper_bound = gr.Slider(minimum=0, maximum=255, value=high_int, step=1, label="Upper Intensity Bound (0-255)")
param1 = gr.Slider(minimum=10, maximum=200, value=edge_thresh, step=1, label="Hough Param1 (Edge Threshold)")
param2 = gr.Slider(minimum=1, maximum=50, value=accum_thresh, step=1, label="Hough Param2 (Accumulator Threshold)")
center_tolerance = gr.Slider(minimum=10, maximum=100, value=center_tol, step=1, label="Center Tolerance (Pixels)")
morph_iterations = gr.Slider(minimum=1, maximum=5, value=morph_dia, step=1, label="Morphological Dilation Iterations")
min_rad = gr.Slider(minimum=1, maximum=100, value=min_rad, step=1, label="Minimum Circle Radius")
display_mode = gr.Dropdown(
choices=["All Frames", "Detected Frames", "Both (Detected Replaces Original)"],
value="Detected Frames",
label="Display Mode"
)
analyze_button = gr.Button("Analyze")
with gr.Column():
gr.Markdown("### Outputs")
report = gr.Textbox(label="Analysis Report", lines=10)
preview = gr.Gallery(label="Input Preview (All Frames)")
gallery = gr.Gallery(label="Frames with Detected Circles (Green: Detected, Orange: Growing Series)")
gif_output = gr.File(label="Download Resulting GIF")
# Fetch button action
fetch_button.click(
fn=handle_fetch,
inputs=[start_date, end_date, ident, size, tool],
outputs=[report, preview, fetched_frames_state]
)
# Analyze button action
analyze_button.click(
fn=process_input,
inputs=[
gif_input, start_date, end_date, ident, size, tool,
lower_bound, upper_bound, param1, param2, center_tolerance, morph_iterations, min_rad, display_mode, fetched_frames_state
],
outputs=[report, gallery, gif_output, preview]
)
if __name__ == "__main__":
demo.launch()