aboutsummaryrefslogtreecommitdiffhomepage
path: root/notebooks/fixed_sankey_plot.py
diff options
context:
space:
mode:
authorAdam Janovsky2023-11-14 17:16:05 +0100
committerAdam Janovsky2023-11-14 17:16:05 +0100
commit0e4a68fd60b21bcfc7f8e27d02c6b9b8ee591f5a (patch)
tree2e431ae258c27b29bb9172b36b810ef6a8af9327 /notebooks/fixed_sankey_plot.py
parentab55c4bd87d83d23b24ada7007c93435230ab158 (diff)
downloadsec-certs-0e4a68fd60b21bcfc7f8e27d02c6b9b8ee591f5a.tar.gz
sec-certs-0e4a68fd60b21bcfc7f8e27d02c6b9b8ee591f5a.tar.zst
sec-certs-0e4a68fd60b21bcfc7f8e27d02c6b9b8ee591f5a.zip
continue refactoring the notebook
Diffstat (limited to 'notebooks/fixed_sankey_plot.py')
-rw-r--r--notebooks/fixed_sankey_plot.py62
1 files changed, 31 insertions, 31 deletions
diff --git a/notebooks/fixed_sankey_plot.py b/notebooks/fixed_sankey_plot.py
index 71c1683a..b8d062f9 100644
--- a/notebooks/fixed_sankey_plot.py
+++ b/notebooks/fixed_sankey_plot.py
@@ -1,5 +1,5 @@
# type: ignore
-
+# ruff: noqa: UP007
"""
This is a fork of https://github.com/anazalea/pySankey/blob/master/pysankey/sankey.py.
We've had some problems with the plot, mostly related to resizing (likely, I don't remember now).
@@ -9,7 +9,7 @@ This code should fix the problems and should be used to produce figures in the r
import logging
import warnings
from collections import defaultdict
-from typing import Any, Dict, List, Optional, Set, Tuple, Union
+from typing import Any, Optional, Union
import matplotlib.pyplot as plt
import numpy as np
@@ -35,7 +35,7 @@ class LabelMismatch(PySankeyException):
LOGGER = logging.getLogger(__name__)
-def check_data_matches_labels(labels: Union[List[str], Set[str]], data: Series, side: str) -> None:
+def check_data_matches_labels(labels: Union[list[str], set[str]], data: Series, side: str) -> None:
"""Check whether data matches labels.
Raise a LabelMismatch Exception if not."""
if len(labels) > 0:
@@ -55,19 +55,19 @@ def check_data_matches_labels(labels: Union[List[str], Set[str]], data: Series,
def sankey(
- left: Union[List, ndarray, Series],
+ left: Union[list, ndarray, Series],
right: Union[ndarray, Series],
leftWeight: Optional[ndarray] = None,
rightWeight: Optional[ndarray] = None,
- colorDict: Optional[Dict[str, str]] = None,
- leftLabels: Optional[List[str]] = None,
- rightLabels: Optional[List[str]] = None,
+ colorDict: Optional[dict[str, str]] = None,
+ leftLabels: Optional[list[str]] = None,
+ rightLabels: Optional[list[str]] = None,
aspect: int = 4,
rightColor: bool = False,
fontsize: int = 14,
figureName: Optional[str] = None,
closePlot: bool = False,
- figSize: Optional[Tuple[int, int]] = None,
+ figSize: Optional[tuple[int, int]] = None,
ax: Optional[Any] = None,
) -> Any:
"""
@@ -158,7 +158,7 @@ def save_image(figureName: Optional[str]) -> None:
LOGGER.info("Sankey diagram generated in '%s'", file_name)
-def identify_labels(dataFrame: DataFrame, leftLabels: List[str], rightLabels: List[str]) -> Tuple[ndarray, ndarray]:
+def identify_labels(dataFrame: DataFrame, leftLabels: list[str], rightLabels: list[str]) -> tuple[ndarray, ndarray]:
# Identify left labels
if len(leftLabels) == 0:
leftLabels = pd.Series(dataFrame.left.unique()).unique()
@@ -175,14 +175,14 @@ def identify_labels(dataFrame: DataFrame, leftLabels: List[str], rightLabels: Li
def init_values(
ax: Optional[Any],
closePlot: bool,
- figSize: Optional[Tuple[int, int]],
+ figSize: Optional[tuple[int, int]],
figureName: Optional[str],
- left: Union[List, ndarray, Series],
- leftLabels: Optional[List[str]],
+ left: Union[list, ndarray, Series],
+ leftLabels: Optional[list[str]],
leftWeight: Optional[ndarray],
- rightLabels: Optional[List[str]],
+ rightLabels: Optional[list[str]],
rightWeight: Optional[ndarray],
-) -> Tuple[Any, List[str], ndarray, List[str], ndarray]:
+) -> tuple[Any, list[str], ndarray, list[str], ndarray]:
deprecation_warnings(closePlot, figSize, figureName)
if ax is None:
ax = plt.gca()
@@ -202,7 +202,7 @@ def init_values(
return ax, leftLabels, leftWeight, rightLabels, rightWeight
-def deprecation_warnings(closePlot: bool, figSize: Optional[Tuple[int, int]], figureName: Optional[str]) -> None:
+def deprecation_warnings(closePlot: bool, figSize: Optional[tuple[int, int]], figureName: Optional[str]) -> None:
warn = []
if figureName is not None:
msg = "use of figureName in sankey() is deprecated"
@@ -223,10 +223,10 @@ def deprecation_warnings(closePlot: bool, figSize: Optional[Tuple[int, int]], fi
)
-def determine_widths(dataFrame: DataFrame, leftLabels: ndarray, rightLabels: ndarray) -> Tuple[Dict, Dict]:
+def determine_widths(dataFrame: DataFrame, leftLabels: ndarray, rightLabels: ndarray) -> tuple[dict, dict]:
# Determine widths of individual strips
- ns_l: Dict = defaultdict()
- ns_r: Dict = defaultdict()
+ ns_l: dict = defaultdict()
+ ns_r: dict = defaultdict()
for leftLabel in leftLabels:
left_dict = {}
right_dict = {}
@@ -244,12 +244,12 @@ def determine_widths(dataFrame: DataFrame, leftLabels: ndarray, rightLabels: nda
def draw_vertical_bars(
ax: Any,
- colorDict: Union[Dict[str, Tuple[float, float, float]], Dict[str, str]],
+ colorDict: Union[dict[str, tuple[float, float, float]], dict[str, str]],
fontsize: int,
leftLabels: ndarray,
- leftWidths: Dict,
+ leftWidths: dict,
rightLabels: ndarray,
- rightWidths: Dict,
+ rightWidths: dict,
xMax: float64,
) -> None:
# Draw vertical bars on left and right of each label's section & print label
@@ -286,8 +286,8 @@ def draw_vertical_bars(
def create_colors(
- allLabels: ndarray, colorDict: Optional[Dict[str, str]]
-) -> Union[Dict[str, Tuple[float, float, float]], Dict[str, str]]:
+ allLabels: ndarray, colorDict: Optional[dict[str, str]]
+) -> Union[dict[str, tuple[float, float, float]], dict[str, str]]:
# If no colorDict given, make one
if colorDict is None:
colorDict = {}
@@ -306,7 +306,7 @@ def create_colors(
def _create_dataframe(
- left: Union[List, ndarray, Series],
+ left: Union[list, ndarray, Series],
leftWeight: Union[ndarray, Series],
right: Union[ndarray, Series],
rightWeight: Union[ndarray, Series],
@@ -336,15 +336,15 @@ def _create_dataframe(
def plot_strips(
ax: Any,
- colorDict: Union[Dict[str, Tuple[float, float, float]], Dict[str, str]],
+ colorDict: Union[dict[str, tuple[float, float, float]], dict[str, str]],
dataFrame: DataFrame,
leftLabels: ndarray,
- leftWidths: Dict,
- ns_l: Dict,
- ns_r: Dict,
+ leftWidths: dict,
+ ns_l: dict,
+ ns_r: dict,
rightColor: bool,
rightLabels: ndarray,
- rightWidths: Dict,
+ rightWidths: dict,
xMax: float64,
) -> None:
# Plot strips
@@ -380,9 +380,9 @@ def plot_strips(
ax.axis("off")
-def _get_positions_and_total_widths(df: DataFrame, labels: ndarray, side: str) -> Tuple[Dict, float64]:
+def _get_positions_and_total_widths(df: DataFrame, labels: ndarray, side: str) -> tuple[dict, float64]:
"""Determine positions of label patches and total widths"""
- widths: Dict = defaultdict()
+ widths: dict = defaultdict()
for i, label in enumerate(labels):
label_widths = {}
label_widths[side] = df[df[side] == label][side + "Weight"].sum()