mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-21 18:50:10 +02:00
Fix typing issues
This commit is contained in:
parent
bbe89e2315
commit
2b582fc2a0
@ -1,5 +1,5 @@
|
||||
from pathlib import Path
|
||||
from typing import List, Literal, Sequence
|
||||
from typing import List, Literal, Sequence, TypedDict
|
||||
from uuid import UUID
|
||||
|
||||
import numpy as np
|
||||
@ -26,6 +26,11 @@ class ParquetOutputConfig(BaseConfig):
|
||||
include_geometry: bool = True
|
||||
|
||||
|
||||
class ClipInfo(TypedDict):
|
||||
clip: data.Clip
|
||||
preds: list[Detection]
|
||||
|
||||
|
||||
class ParquetFormatter(OutputFormatterProtocol[ClipDetections]):
|
||||
def __init__(
|
||||
self,
|
||||
@ -120,7 +125,7 @@ class ParquetFormatter(OutputFormatterProtocol[ClipDetections]):
|
||||
else:
|
||||
df = pd.read_parquet(path)
|
||||
|
||||
predictions_by_clip = {}
|
||||
predictions_by_clip: dict[UUID, ClipInfo] = {}
|
||||
|
||||
for _, row in df.iterrows():
|
||||
clip_uuid = row["clip_uuid"]
|
||||
|
||||
@ -185,25 +185,25 @@ def plot_clip_evaluation(
|
||||
label="found GT",
|
||||
edgecolor=gt_color,
|
||||
facecolor="none" if not fill else gt_color,
|
||||
linestyle=gt_linestyle,
|
||||
linestyle=gt_linestyle, # type: ignore
|
||||
),
|
||||
patches.Patch(
|
||||
label="missed GT",
|
||||
edgecolor=missed_gt_color,
|
||||
facecolor="none" if not fill else missed_gt_color,
|
||||
linestyle=missed_gt_linestyle,
|
||||
linestyle=missed_gt_linestyle, # type: ignore
|
||||
),
|
||||
patches.Patch(
|
||||
label="true Det",
|
||||
edgecolor=true_pred_color,
|
||||
facecolor="none" if not fill else true_pred_color,
|
||||
linestyle=true_pred_linestyle,
|
||||
linestyle=true_pred_linestyle, # type: ignore
|
||||
),
|
||||
patches.Patch(
|
||||
label="false Det",
|
||||
edgecolor=false_pred_color,
|
||||
facecolor="none" if not fill else false_pred_color,
|
||||
linestyle=false_pred_linestyle,
|
||||
linestyle=false_pred_linestyle, # type: ignore
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@ -1,9 +1,9 @@
|
||||
"""Plot heatmaps."""
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from matplotlib import axes, patches
|
||||
from matplotlib.cm import get_cmap
|
||||
from matplotlib.colors import Colormap, LinearSegmentedColormap, to_rgba
|
||||
|
||||
from batdetect2.plotting.common import create_ax
|
||||
@ -80,7 +80,7 @@ def plot_classification_heatmap(
|
||||
raise ValueError("Inconsistent number of class names")
|
||||
|
||||
if not isinstance(cmap, Colormap):
|
||||
cmap = get_cmap(cmap)
|
||||
cmap = plt.get_cmap(cmap)
|
||||
|
||||
handles = []
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user