[PlotMetricCurvesPIL] Multiple curves and more
- Multiple traces: colors and legend_labels - auto_scale option - Better grid division - Margin, font, trace size auto-adjust
This commit is contained in:
+334
-98
@@ -1183,25 +1183,29 @@ class ConsolidateMetrics(ComfyNodeABC):
|
||||
return ([v for v in existing_data.values()], )
|
||||
|
||||
|
||||
# Most code from Gemini 3 Pro
|
||||
class PlotMetricCurvesPIL(ComfyNodeABC):
|
||||
"""
|
||||
Plots the Precision vs Recall and F-measure curves
|
||||
|
||||
Generates two images containing the curves from the consolidated `metrics`
|
||||
"""
|
||||
# Define available colors for the plot line
|
||||
COLORS = ['blue', 'green', 'red', 'cyan', 'magenta', 'black']
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"metrics": ("DICT",),
|
||||
"plot_title": (IO.STRING, {"default": "Saliency Evaluation"}),
|
||||
"curve_color": (s.COLORS,),
|
||||
"curve_color": (IO.STRING, {"default": "red"}),
|
||||
"width": (IO.INT, {"default": 800, "min": 256, "max": 4096}),
|
||||
"height": (IO.INT, {"default": 600, "min": 256, "max": 4096}),
|
||||
"auto_scale": (IO.BOOLEAN, {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"legend_labels": ("STRING", {
|
||||
"forceInput": True,
|
||||
"tooltip": "List of names for the curves. If connected, a legend is drawn."}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
@@ -1217,158 +1221,390 @@ class PlotMetricCurvesPIL(ComfyNodeABC):
|
||||
# Convert to numpy array, normalize to [0, 1], and add batch dimension
|
||||
return torch.from_numpy(np.array(pil_image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
def _get_nice_limits(self, data_min, data_max, padding_percent=0.05):
|
||||
"""Calculates 'nice' axis limits with padding, handling edge cases."""
|
||||
# Handle the case where all data points are the same
|
||||
if data_min == data_max:
|
||||
return data_min - 0.1, data_max + 0.1
|
||||
|
||||
# Calculate padding based on the data range
|
||||
data_range = data_max - data_min
|
||||
padding = data_range * padding_percent
|
||||
|
||||
# Return the padded limits
|
||||
return data_min - padding, data_max + padding
|
||||
|
||||
def _create_plot_with_pil(self, data_x, data_y, title, x_label, y_label, width, height, color, x_lim=None, y_lim=None):
|
||||
def _get_nice_ticks(self, min_val, max_val, max_ticks=12):
|
||||
"""
|
||||
Generates a plot from scratch using PIL with auto-scaling and rotated Y-axis text.
|
||||
Calculates 'nice' tick locations for a graph axis, similar to matplotlib.
|
||||
Returns (tick_values, view_min, view_max).
|
||||
"""
|
||||
# --- 1. Auto-scale limits if they are not provided ---
|
||||
# 1. Basic range check
|
||||
if min_val == max_val:
|
||||
return [min_val], min_val - 0.1, max_val + 0.1
|
||||
|
||||
raw_range = max_val - min_val
|
||||
|
||||
# 2. Calculate a rough step size
|
||||
rough_step = raw_range / (max_ticks - 1)
|
||||
|
||||
# 3. Calculate the 'magnitude' of the step (power of 10)
|
||||
# E.g., if rough_step is 55, magnitude is 10. If 0.05, magnitude is 0.01
|
||||
magnitude = 10 ** np.floor(np.log10(rough_step))
|
||||
|
||||
# 4. Normalize the rough step to [1, 10)
|
||||
normalized_step = rough_step / magnitude
|
||||
|
||||
# 5. Pick the closest 'nice' step multiple (1, 2, 5, 10)
|
||||
if normalized_step <= 1.0:
|
||||
nice_step = 1.0 * magnitude
|
||||
elif normalized_step <= 2.0:
|
||||
nice_step = 2.0 * magnitude
|
||||
elif normalized_step <= 5.0:
|
||||
nice_step = 5.0 * magnitude
|
||||
else:
|
||||
nice_step = 10.0 * magnitude
|
||||
|
||||
# 6. Calculate new min/max values based on the nice step
|
||||
# We floor the min and ceil the max to the nearest nice step
|
||||
tick_min = np.floor(min_val / nice_step) * nice_step
|
||||
tick_max = np.ceil(max_val / nice_step) * nice_step
|
||||
|
||||
# 7. Generate the tick values
|
||||
# We use np.arange but being careful with float precision
|
||||
# Adding half a step to the stop value ensures the last tick is included
|
||||
ticks = np.arange(tick_min, tick_max + nice_step/2, nice_step)
|
||||
|
||||
# Filter ticks to keep them reasonably close to data (optional but cleaner)
|
||||
# Here we strictly respect the calculated nice bounds
|
||||
return ticks, tick_min, tick_max
|
||||
|
||||
def _create_plot_with_pil(self, curves_data, title, x_label, y_label, width, height, x_lim=None, y_lim=None):
|
||||
"""
|
||||
Generates a plot from scratch using PIL.
|
||||
|
||||
Args:
|
||||
curves_data: A list of tuples (data_x, data_y, color_string)
|
||||
"""
|
||||
# --- Determine Data Ranges ---
|
||||
if x_lim is None:
|
||||
x_lim = self._get_nice_limits(np.min(data_x), np.max(data_x))
|
||||
# Flatten all X data to find global min/max
|
||||
all_x = [val for c in curves_data for val in c[0]]
|
||||
data_x_min, data_x_max = np.min(all_x), np.max(all_x)
|
||||
else:
|
||||
data_x_min, data_x_max = x_lim
|
||||
if y_lim is None:
|
||||
y_lim = self._get_nice_limits(np.min(data_y), np.max(data_y))
|
||||
# Flatten all Y data to find global min/max
|
||||
all_y = [val for c in curves_data for val in c[1]]
|
||||
data_y_min, data_y_max = np.min(all_y), np.max(all_y)
|
||||
else:
|
||||
data_y_min, data_y_max = y_lim
|
||||
|
||||
# --- 2. Setup Canvas and Drawing Tools ---
|
||||
padding_left = 80 # Increased padding to accommodate rotated label
|
||||
padding_right = 30
|
||||
padding_top = 60
|
||||
padding_bottom = 60
|
||||
# --- Calculate 'Nice' Ticks and View Limits ---
|
||||
x_ticks, view_x_min, view_x_max = self._get_nice_ticks(data_x_min, data_x_max)
|
||||
y_ticks, view_y_min, view_y_max = self._get_nice_ticks(data_y_min, data_y_max)
|
||||
|
||||
# If hard limits were passed (like 0.0 to 1.0), we might want to clamp the view
|
||||
# but keep the nice internal ticks.
|
||||
if x_lim is not None:
|
||||
view_x_min, view_x_max = x_lim
|
||||
# Filter ticks outside limits
|
||||
x_ticks = [t for t in x_ticks if t >= view_x_min and t <= view_x_max]
|
||||
|
||||
if y_lim is not None:
|
||||
view_y_min, view_y_max = y_lim
|
||||
y_ticks = [t for t in y_ticks if t >= view_y_min and t <= view_y_max]
|
||||
|
||||
# --- Setup Canvas and Drawing Tools ---
|
||||
# Base scale on the smaller dimension to ensure fit
|
||||
base_dim = min(width, height)
|
||||
|
||||
# Calculate sizes based on resolution
|
||||
# Heuristics:
|
||||
# - Title: ~1/25th of height
|
||||
# - Axis Labels: ~1/35th of height
|
||||
# - Tick Labels: ~1/45th of height
|
||||
# - Trace Width: ~1/150th of height
|
||||
title_size = max(16, int(base_dim / 20))
|
||||
axis_label_size = max(12, int(base_dim / 20))
|
||||
tick_label_size = max(10, int(base_dim / 30))
|
||||
|
||||
trace_width = max(2, int(base_dim / 150))
|
||||
grid_width = max(1, int(trace_width / 3))
|
||||
frame_width = max(1, int(trace_width / 2))
|
||||
|
||||
# Dynamic Padding based on font sizes
|
||||
# We need room for the Y labels (left), Title (top), X labels (bottom)
|
||||
padding_top = int(title_size * 2.5)
|
||||
padding_bottom = int(axis_label_size + tick_label_size * 2.5)
|
||||
padding_left = int(axis_label_size + tick_label_size * 4.0) # Room for "0.00" + axis label
|
||||
padding_right = int(tick_label_size * 2)
|
||||
|
||||
img = Image.new('RGB', (width, height), 'white')
|
||||
draw = ImageDraw.Draw(img)
|
||||
font = load_font("Arial", 15)
|
||||
title_font = load_font("Arial", 20)
|
||||
title_font = load_font("Arial", title_size)
|
||||
axis_font = load_font("Arial", axis_label_size)
|
||||
tick_font = load_font("Arial", tick_label_size)
|
||||
|
||||
# --- 3. Define Plot Area and Coordinate Mapping ---
|
||||
# (This section is unchanged)
|
||||
# --- Define Plot Area and Coordinate Mapping ---
|
||||
plot_width = width - padding_left - padding_right
|
||||
plot_height = height - padding_top - padding_bottom
|
||||
|
||||
# Coordinate mapping using the VIEW limits (calculated by nice ticks)
|
||||
def to_pixel(x, y):
|
||||
px = padding_left + ((x - x_lim[0]) / (x_lim[1] - x_lim[0])) * plot_width
|
||||
py = (height - padding_bottom) - ((y - y_lim[0]) / (y_lim[1] - y_lim[0])) * plot_height
|
||||
# Normalize to 0-1 range relative to view
|
||||
x_norm = (x - view_x_min) / (view_x_max - view_x_min) if view_x_max > view_x_min else 0.5
|
||||
y_norm = (y - view_y_min) / (view_y_max - view_y_min) if view_y_max > view_y_min else 0.5
|
||||
|
||||
# Clamp for safety (don't draw way off canvas)
|
||||
x_norm = max(0.0, min(1.0, x_norm))
|
||||
y_norm = max(0.0, min(1.0, y_norm))
|
||||
|
||||
px = padding_left + x_norm * plot_width
|
||||
py = (height - padding_bottom) - y_norm * plot_height
|
||||
return int(px), int(py)
|
||||
|
||||
# --- 4. Draw Grid, Axes, and Ticks ---
|
||||
# (This section is unchanged)
|
||||
num_grid_lines = 5
|
||||
for i in range(num_grid_lines + 1):
|
||||
val_x = x_lim[0] + (i / num_grid_lines) * (x_lim[1] - x_lim[0])
|
||||
px, _ = to_pixel(val_x, y_lim[0])
|
||||
draw.line([(px, padding_top), (px, height - padding_bottom)], fill=(220, 220, 220), width=1)
|
||||
label = f"{val_x:.2f}" if val_x % 1 else str(int(val_x))
|
||||
draw.text((px, height - padding_bottom + 5), label, font=font, fill='black', anchor="mt")
|
||||
# --- Draw Grid, Axes, and Ticks ---
|
||||
# X-Axis Ticks
|
||||
for val in x_ticks:
|
||||
px, _ = to_pixel(val, view_y_min)
|
||||
|
||||
for i in range(num_grid_lines + 1):
|
||||
val_y = y_lim[0] + (i / num_grid_lines) * (y_lim[1] - y_lim[0])
|
||||
_, py = to_pixel(x_lim[0], val_y)
|
||||
draw.line([(padding_left, py), (width - padding_right, py)], fill=(220, 220, 220), width=1)
|
||||
label = f"{val_y:.2f}"
|
||||
draw.text((padding_left - 10, py), label, font=font, fill='black', anchor="rm")
|
||||
# Grid line
|
||||
draw.line([(px, padding_top), (px, height - padding_bottom)], fill=(230, 230, 230), width=grid_width)
|
||||
|
||||
draw.line([(padding_left, height - padding_bottom), (width - padding_right, height - padding_bottom)], fill='black',
|
||||
width=2)
|
||||
draw.line([(padding_left, padding_top), (padding_left, height - padding_bottom)], fill='black', width=2)
|
||||
# Tick label
|
||||
# Smart formatting: remove decimal if integer
|
||||
if abs(val - round(val)) < 1e-8:
|
||||
label = f"{int(round(val))}"
|
||||
else:
|
||||
label = f"{val:.2f}".rstrip('0').rstrip('.')
|
||||
|
||||
# --- 5. Draw the Data Curve ---
|
||||
# (This section is unchanged)
|
||||
pixel_points = [to_pixel(x, y) for x, y in zip(data_x, data_y) if x_lim[0] <= x <= x_lim[1] and
|
||||
y_lim[0] <= y <= y_lim[1]]
|
||||
if len(pixel_points) > 1:
|
||||
draw.line(pixel_points, fill=color, width=3)
|
||||
draw.text((px, height - padding_bottom + 8), label, font=tick_font, fill='black', anchor="mt")
|
||||
|
||||
# --- 6. Draw Title and Labels ---
|
||||
# Y-Axis Ticks
|
||||
for val in y_ticks:
|
||||
_, py = to_pixel(view_x_min, val)
|
||||
|
||||
# Draw Title and X-axis Label (unchanged)
|
||||
# Grid line
|
||||
draw.line([(padding_left, py), (width - padding_right, py)], fill=(230, 230, 230), width=grid_width)
|
||||
|
||||
# Tick label
|
||||
if abs(val - round(val)) < 1e-8:
|
||||
label = f"{int(round(val))}"
|
||||
else:
|
||||
label = f"{val:.2f}".rstrip('0').rstrip('.')
|
||||
|
||||
draw.text((padding_left - 10, py), label, font=tick_font, fill='black', anchor="rm")
|
||||
|
||||
# Draw Axis Frames
|
||||
draw.rectangle([padding_left, padding_top, width - padding_right, height - padding_bottom], outline='black',
|
||||
width=frame_width)
|
||||
|
||||
# --- Draw the Data Curve ---
|
||||
for data_x, data_y, color, label in curves_data:
|
||||
# Map all points
|
||||
pixel_points = []
|
||||
for x, y in zip(data_x, data_y):
|
||||
# Only draw if within view range (plus a tiny epsilon margin)
|
||||
if x >= view_x_min and x <= view_x_max and y >= view_y_min and y <= view_y_max:
|
||||
pixel_points.append(to_pixel(x, y))
|
||||
|
||||
if len(pixel_points) > 1:
|
||||
draw.line(pixel_points, fill=color, width=trace_width)
|
||||
|
||||
# --- Draw Legend ---
|
||||
# Check if we have labels. If the first item has a label, we assume we draw legends.
|
||||
if curves_data and curves_data[0][3] is not None:
|
||||
|
||||
# Legend Style Settings
|
||||
leg_padding = int(tick_label_size) # Internal padding
|
||||
leg_line_len = int(base_dim / 25) # Length of the colored line sample
|
||||
leg_gap = int(base_dim / 60) # Gap between line and text
|
||||
line_height = int(tick_label_size * 1.4) # Height of one legend row
|
||||
|
||||
# Calculate Legend Box Size
|
||||
max_text_width = 0
|
||||
for _, _, _, label in curves_data:
|
||||
bbox = tick_font.getbbox(str(label))
|
||||
w = bbox[2] - bbox[0]
|
||||
if w > max_text_width:
|
||||
max_text_width = w
|
||||
|
||||
box_width = leg_padding * 2 + leg_line_len + leg_gap + max_text_width
|
||||
box_height = leg_padding * 2 + (len(curves_data) * line_height)
|
||||
|
||||
# Position: Lower Left (inside plot area)
|
||||
box_x = padding_left + 15
|
||||
box_y = (height - padding_bottom) - box_height - 15
|
||||
|
||||
# Draw Semi-Transparent Background
|
||||
# PIL requires RGBA mode for alpha compositing
|
||||
overlay = Image.new('RGBA', img.size, (0, 0, 0, 0))
|
||||
draw_overlay = ImageDraw.Draw(overlay)
|
||||
|
||||
# Background (White with ~75% opacity -> 190 alpha)
|
||||
draw_overlay.rectangle(
|
||||
[box_x, box_y, box_x + box_width, box_y + box_height],
|
||||
fill=(255, 255, 255, 190),
|
||||
outline=(100, 100, 100, 255), # Dark gray border
|
||||
width=1
|
||||
)
|
||||
|
||||
# Composite overlay onto main image
|
||||
img = img.convert('RGBA')
|
||||
img = Image.alpha_composite(img, overlay)
|
||||
img = img.convert('RGB')
|
||||
draw = ImageDraw.Draw(img) # Re-acquire draw object for RGB
|
||||
|
||||
# Draw Legend Items
|
||||
curr_y = box_y + leg_padding
|
||||
line_y_offset = line_height // 2
|
||||
|
||||
for _, _, color, label in curves_data:
|
||||
# Draw colored line sample
|
||||
line_start_x = box_x + leg_padding
|
||||
line_end_x = line_start_x + leg_line_len
|
||||
line_y = curr_y + line_y_offset
|
||||
draw.line([(line_start_x, line_y), (line_end_x, line_y)], fill=color, width=max(2, trace_width))
|
||||
|
||||
# Draw text
|
||||
text_x = line_end_x + leg_gap
|
||||
# Adjust text_y to vertically center with the line
|
||||
# tick_font.getbbox returns (left, top, right, bottom) relative to baseline
|
||||
# simple approach: slightly nudge up based on font size
|
||||
text_y = curr_y + (line_height - tick_label_size) / 2
|
||||
|
||||
draw.text((text_x, text_y), str(label), font=tick_font, fill='black')
|
||||
|
||||
curr_y += line_height
|
||||
|
||||
# --- Draw Title and Labels ---
|
||||
# Draw Title and X-axis Label
|
||||
draw.text((width / 2, padding_top / 2), title, font=title_font, fill='black', anchor="mm")
|
||||
draw.text((width / 2, height - padding_bottom / 4), x_label, font=title_font, fill='black', anchor="mb")
|
||||
draw.text((width / 2, height - padding_bottom / 4), x_label, font=axis_font, fill='black', anchor="mb")
|
||||
|
||||
# --- Draw Rotated Y-axis Label ---
|
||||
|
||||
# a. Get the size of the unrotated text
|
||||
# y_label_bbox = font.getbbox(y_label)
|
||||
# y_label_width = y_label_bbox[2] - y_label_bbox[0]
|
||||
# y_label_height = y_label_bbox[3] - y_label_bbox[1]
|
||||
|
||||
# b. Create a new, transparent canvas for the text
|
||||
# txt_canvas = Image.new('RGBA', (y_label_width, y_label_height), (0, 0, 0, 0))
|
||||
# Create a new, transparent canvas for the text
|
||||
txt_canvas = Image.new('RGBA', (height, padding_left), (0, 0, 0, 0))
|
||||
txt_draw = ImageDraw.Draw(txt_canvas)
|
||||
|
||||
# c. Draw the text onto the temporary canvas
|
||||
txt_draw.text((height // 2, padding_left // 4), y_label, font=title_font, fill='black', anchor="mb")
|
||||
# Draw the text onto the temporary canvas
|
||||
txt_draw.text((height // 2, axis_label_size), y_label, font=axis_font, fill='black', anchor="mm")
|
||||
|
||||
# d. Rotate the text canvas by 90 degrees
|
||||
# Rotate the text canvas by 90 degrees
|
||||
# 'expand=True' makes the new image large enough to hold the rotated content
|
||||
rotated_y_label = txt_canvas.rotate(90, expand=True)
|
||||
|
||||
# e. Calculate the paste position on the main canvas
|
||||
# Calculate the paste position on the main canvas
|
||||
# Center it vertically in the plot area and horizontally in the left padding area
|
||||
paste_x = 0 # int((padding_left - rotated_y_label.width) / 2)
|
||||
paste_x = 0
|
||||
paste_y = int((height - rotated_y_label.height) / 2)
|
||||
|
||||
# f. Paste the rotated text onto the main image, using its alpha channel as a mask
|
||||
# Paste the rotated text onto the main image, using its alpha channel as a mask
|
||||
img.paste(rotated_y_label, (paste_x, paste_y), rotated_y_label)
|
||||
|
||||
return img
|
||||
|
||||
def execute(self, metrics, plot_title, curve_color, width, height):
|
||||
# --- 1. Aggregate Data (same as matplotlib version) ---
|
||||
def execute(self, metrics, plot_title, curve_color, width, height, auto_scale, legend_labels=None):
|
||||
# --- Aggregate Data (same as matplotlib version) ---
|
||||
plot_title_str = plot_title[0]
|
||||
curve_color_str = curve_color[0]
|
||||
color_str = curve_color[0]
|
||||
auto_scale = auto_scale[0]
|
||||
|
||||
if not metrics:
|
||||
logger.warning("No metrics data provided. Returning blank images.")
|
||||
blank_image = self._pil_to_tensor(Image.new('RGB', (width[0], height[0]), 'white'))
|
||||
return (blank_image, blank_image)
|
||||
|
||||
all_precisions, all_recalls, all_fmeasures = [], [], []
|
||||
# --- Normalize Input to List of Datasets ---
|
||||
# Check if the input is a single dataset (List of Dicts) or multiple (List of Lists of Dicts)
|
||||
# We assume if the first item is a dict, it's a single dataset.
|
||||
if isinstance(metrics[0], dict):
|
||||
datasets = [metrics]
|
||||
else:
|
||||
datasets = metrics
|
||||
|
||||
for metric_dict in metrics:
|
||||
# Check if the dictionary contains the compressed tensor keys.
|
||||
if 'fp' in metric_dict and 'fr' in metric_dict and 'f' in metric_dict:
|
||||
# Retrieve the tensor, move to CPU, and convert to a NumPy array.
|
||||
# The .cpu() is important for safety in case tensors are on the GPU.
|
||||
all_precisions.append(metric_dict['fp'].cpu().numpy())
|
||||
all_recalls.append(metric_dict['fr'].cpu().numpy())
|
||||
all_fmeasures.append(metric_dict['f'].cpu().numpy())
|
||||
# --- Prepare Colors ---
|
||||
colors = color_str.split()
|
||||
n_colors = len(colors)
|
||||
n_datasets = len(datasets)
|
||||
if n_colors != n_datasets:
|
||||
logger.warning(f"Got {n_datasets} curves and {n_colors} colors")
|
||||
if n_colors < n_datasets:
|
||||
logger.debug("Repeating the last color")
|
||||
colors += [colors[-1]] * (n_datasets - n_colors)
|
||||
|
||||
if not all_precisions:
|
||||
# --- Prepare Labels ---
|
||||
labels = None
|
||||
if legend_labels is not None:
|
||||
# Ensure it's a list and has enough entries
|
||||
# legend_labels comes in as a list of strings usually if from a primitive
|
||||
if isinstance(legend_labels, list):
|
||||
labels = legend_labels
|
||||
else:
|
||||
labels = [legend_labels] # Handle single string edge case
|
||||
|
||||
# Pad labels if fewer than datasets
|
||||
if len(labels) < n_datasets:
|
||||
labels += [f"Curve {i+1}" for i in range(len(labels), n_datasets)]
|
||||
|
||||
# --- Process Each Dataset ---
|
||||
pr_curves_data = [] # Will hold tuples: (x_data, y_data, color)
|
||||
fm_curves_data = [] # Will hold tuples: (x_data, y_data, color)
|
||||
|
||||
# Track global min/max for scaling
|
||||
global_pr_x_min = float('inf')
|
||||
global_pr_x_max = float('-inf')
|
||||
|
||||
for i, dataset in enumerate(datasets):
|
||||
all_precisions, all_recalls, all_fmeasures = [], [], []
|
||||
|
||||
for metric_dict in dataset:
|
||||
# Check if the dictionary contains the compressed tensor keys.
|
||||
if 'fp' in metric_dict and 'fr' in metric_dict and 'f' in metric_dict:
|
||||
# Retrieve the tensor, move to CPU, and convert to a NumPy array.
|
||||
# The .cpu() is important for safety in case tensors are on the GPU.
|
||||
all_precisions.append(metric_dict['fp'].cpu().numpy())
|
||||
all_recalls.append(metric_dict['fr'].cpu().numpy())
|
||||
all_fmeasures.append(metric_dict['f'].cpu().numpy())
|
||||
|
||||
if not all_precisions:
|
||||
continue # Skip empty datasets
|
||||
|
||||
avg_precision = np.mean(all_precisions, axis=0)
|
||||
avg_recall = np.mean(all_recalls, axis=0)
|
||||
avg_fmeasure = np.mean(all_fmeasures, axis=0)
|
||||
|
||||
# Update global stats for PR curve (we ignore Y limits for PR as per your request/standard)
|
||||
global_pr_x_min = min(global_pr_x_min, np.min(avg_recall))
|
||||
global_pr_x_max = max(global_pr_x_max, np.max(avg_recall))
|
||||
|
||||
# Determine label for this curve
|
||||
curr_label = labels[i] if labels else None
|
||||
|
||||
# Store data for plotting
|
||||
pr_curves_data.append((avg_recall, avg_precision, colors[i], curr_label))
|
||||
|
||||
threshold_axis = np.arange(len(avg_fmeasure))
|
||||
fm_curves_data.append((threshold_axis, avg_fmeasure, colors[i], curr_label))
|
||||
|
||||
if not pr_curves_data:
|
||||
logger.warning("Metrics data did not contain F/FP/FR keys. Returning blank images.")
|
||||
blank_image = self._pil_to_tensor(Image.new('RGB', (width[0], height[0]), 'white'))
|
||||
return (blank_image, blank_image)
|
||||
|
||||
avg_precision = np.mean(all_precisions, axis=0)
|
||||
avg_recall = np.mean(all_recalls, axis=0)
|
||||
avg_fmeasure = np.mean(all_fmeasures, axis=0)
|
||||
# Determine limits
|
||||
if auto_scale:
|
||||
# Use the global min/max we calculated
|
||||
pr_x_lim = (global_pr_x_min, global_pr_x_max)
|
||||
# For PR, Y is usually 0-1, or we can auto-scale that too.
|
||||
pr_y_lim = None
|
||||
else:
|
||||
pr_x_lim = (0.0, 1.0)
|
||||
pr_y_lim = (0.0, 1.0)
|
||||
|
||||
# --- 2. Generate Precision-Recall (PR) Curve Plot ---
|
||||
# --- Generate Precision-Recall (PR) Curve Plot ---
|
||||
pr_plot_pil = self._create_plot_with_pil(
|
||||
data_x=avg_recall, data_y=avg_precision,
|
||||
curves_data=pr_curves_data,
|
||||
title=f"{plot_title_str} (PR Curve)", x_label="Recall", y_label="Precision",
|
||||
width=width[0], height=height[0], color=curve_color_str,
|
||||
x_lim=(np.min(avg_recall), np.max(avg_recall)), y_lim=None # (np.min(avg_precision), np.max(avg_precision)
|
||||
width=width[0], height=height[0],
|
||||
x_lim=pr_x_lim, y_lim=pr_y_lim
|
||||
)
|
||||
pr_plot_tensor = self._pil_to_tensor(pr_plot_pil)
|
||||
|
||||
# --- 3. Generate F-Measure Curve Plot ---
|
||||
threshold_axis = np.arange(F_POINTS)
|
||||
# --- Generate F-Measure Curve Plot ---
|
||||
# F-measure is always 0-255 on X and 0-1 on Y
|
||||
fm_plot_pil = self._create_plot_with_pil(
|
||||
data_x=threshold_axis, data_y=avg_fmeasure,
|
||||
curves_data=fm_curves_data,
|
||||
title=f"{plot_title_str} (F-Measure Curve)", x_label="Threshold", y_label="F-measure",
|
||||
width=width[0], height=height[0], color=curve_color_str,
|
||||
width=width[0], height=height[0],
|
||||
x_lim=(0, F_POINTS), y_lim=(0.0, 1.0)
|
||||
)
|
||||
fm_plot_tensor = self._pil_to_tensor(fm_plot_pil)
|
||||
|
||||
Reference in New Issue
Block a user