summaryrefslogtreecommitdiff
path: root/networkx/drawing/layout.py
diff options
context:
space:
mode:
authorDimitrios Papageorgiou <dim_papag@windowslive.com>2021-03-09 07:19:46 +0200
committerGitHub <noreply@github.com>2021-03-08 21:19:46 -0800
commit6ff62cd16317cd3409f8e51fe7abcf3888ff5517 (patch)
treeae91fcf2216c99b3b0cadebc104701785cc5f5d6 /networkx/drawing/layout.py
parent1de57a8ca852cf27820512720968c6cffb51a74b (diff)
downloadnetworkx-6ff62cd16317cd3409f8e51fe7abcf3888ff5517.tar.gz
Refactor bipartite and multipartite layout (#4653)
Reduce duplication in source code
Diffstat (limited to 'networkx/drawing/layout.py')
-rw-r--r--networkx/drawing/layout.py96
1 files changed, 35 insertions, 61 deletions
diff --git a/networkx/drawing/layout.py b/networkx/drawing/layout.py
index 9ad5948f..474556f0 100644
--- a/networkx/drawing/layout.py
+++ b/networkx/drawing/layout.py
@@ -310,6 +310,10 @@ def bipartite_layout(
import numpy as np
+ if align not in ("vertical", "horizontal"):
+ msg = "align must be either vertical or horizontal."
+ raise ValueError(msg)
+
G, center = _process_params(G, center=center, dim=2)
if len(G) == 0:
return {}
@@ -322,36 +326,20 @@ def bipartite_layout(
bottom = set(G) - top
nodes = list(top) + list(bottom)
- if align == "vertical":
- left_xs = np.repeat(0, len(top))
- right_xs = np.repeat(width, len(bottom))
- left_ys = np.linspace(0, height, len(top))
- right_ys = np.linspace(0, height, len(bottom))
+ left_xs = np.repeat(0, len(top))
+ right_xs = np.repeat(width, len(bottom))
+ left_ys = np.linspace(0, height, len(top))
+ right_ys = np.linspace(0, height, len(bottom))
- top_pos = np.column_stack([left_xs, left_ys]) - offset
- bottom_pos = np.column_stack([right_xs, right_ys]) - offset
-
- pos = np.concatenate([top_pos, bottom_pos])
- pos = rescale_layout(pos, scale=scale) + center
- pos = dict(zip(nodes, pos))
- return pos
+ top_pos = np.column_stack([left_xs, left_ys]) - offset
+ bottom_pos = np.column_stack([right_xs, right_ys]) - offset
+ pos = np.concatenate([top_pos, bottom_pos])
+ pos = rescale_layout(pos, scale=scale) + center
if align == "horizontal":
- top_ys = np.repeat(height, len(top))
- bottom_ys = np.repeat(0, len(bottom))
- top_xs = np.linspace(0, width, len(top))
- bottom_xs = np.linspace(0, width, len(bottom))
-
- top_pos = np.column_stack([top_xs, top_ys]) - offset
- bottom_pos = np.column_stack([bottom_xs, bottom_ys]) - offset
-
- pos = np.concatenate([top_pos, bottom_pos])
- pos = rescale_layout(pos, scale=scale) + center
- pos = dict(zip(nodes, pos))
- return pos
-
- msg = "align must be either vertical or horizontal."
- raise ValueError(msg)
+ pos = np.flip(pos, 1)
+ pos = dict(zip(nodes, pos))
+ return pos
@random_state(10)
@@ -1081,6 +1069,10 @@ def multipartite_layout(G, subset_key="subset", align="vertical", scale=1, cente
"""
import numpy as np
+ if align not in ("vertical", "horizontal"):
+ msg = "align must be either vertical or horizontal."
+ raise ValueError(msg)
+
G, center = _process_params(G, center=center, dim=2)
if len(G) == 0:
return {}
@@ -1096,42 +1088,24 @@ def multipartite_layout(G, subset_key="subset", align="vertical", scale=1, cente
pos = None
nodes = []
- if align == "vertical":
- width = len(layers)
- for i, layer in layers.items():
- height = len(layer)
- xs = np.repeat(i, height)
- ys = np.arange(0, height, dtype=float)
- offset = ((width - 1) / 2, (height - 1) / 2)
- layer_pos = np.column_stack([xs, ys]) - offset
- if pos is None:
- pos = layer_pos
- else:
- pos = np.concatenate([pos, layer_pos])
- nodes.extend(layer)
- pos = rescale_layout(pos, scale=scale) + center
- pos = dict(zip(nodes, pos))
- return pos
+ width = len(layers)
+ for i, layer in layers.items():
+ height = len(layer)
+ xs = np.repeat(i, height)
+ ys = np.arange(0, height, dtype=float)
+ offset = ((width - 1) / 2, (height - 1) / 2)
+ layer_pos = np.column_stack([xs, ys]) - offset
+ if pos is None:
+ pos = layer_pos
+ else:
+ pos = np.concatenate([pos, layer_pos])
+ nodes.extend(layer)
+ pos = rescale_layout(pos, scale=scale) + center
if align == "horizontal":
- height = len(layers)
- for i, layer in layers.items():
- width = len(layer)
- xs = np.arange(0, width, dtype=float)
- ys = np.repeat(i, width)
- offset = ((width - 1) / 2, (height - 1) / 2)
- layer_pos = np.column_stack([xs, ys]) - offset
- if pos is None:
- pos = layer_pos
- else:
- pos = np.concatenate([pos, layer_pos])
- nodes.extend(layer)
- pos = rescale_layout(pos, scale=scale) + center
- pos = dict(zip(nodes, pos))
- return pos
-
- msg = "align must be either vertical or horizontal."
- raise ValueError(msg)
+ pos = np.flip(pos, 1)
+ pos = dict(zip(nodes, pos))
+ return pos
def rescale_layout(pos, scale=1):