Skip to content

Commit a613b54

Browse files
committed
Final improvements
take it or leave it ^^
1 parent 473b779 commit a613b54

4 files changed

Lines changed: 68 additions & 10 deletions

File tree

pyglotaran_extras/inspect/kinetic_scheme/_k_matrix_parser.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -177,9 +177,13 @@ def extract_dataset_transitions(
177177

178178
# Filter to only decay-type megacomplexes (silently skip non-decay).
179179
# Use the raw megacomplex object to avoid filling each item twice.
180-
decay_megacomplexes = [
181-
mc for mc in megacomplexes if hasattr(model.megacomplex[mc], "get_k_matrix")
182-
]
180+
decay_megacomplexes: list[str] = []
181+
for mc in megacomplexes:
182+
if mc not in model.megacomplex:
183+
msg = f"Megacomplex '{mc}' referenced by dataset '{dataset_name}' not found in model."
184+
raise ValueError(msg)
185+
if hasattr(model.megacomplex[mc], "get_k_matrix"):
186+
decay_megacomplexes.append(mc)
183187

184188
return extract_transitions(
185189
decay_megacomplexes,

pyglotaran_extras/inspect/kinetic_scheme/_layout.py

Lines changed: 36 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
from __future__ import annotations
44

5+
import hashlib
56
from collections import deque
67
from enum import Enum
78
from typing import TYPE_CHECKING
@@ -401,11 +402,16 @@ def _build_in_degrees(
401402
In-degree for each compartment node.
402403
"""
403404
in_degree: dict[str, int] = dict.fromkeys(compartment_labels, 0)
405+
seen_edges: set[tuple[str, str]] = set()
404406
for edge in graph.edges:
407+
edge_pair = (edge.source, edge.target)
408+
if edge_pair in seen_edges:
409+
continue
410+
seen_edges.add(edge_pair)
405411
if (
406412
edge.source in compartment_labels
407413
and edge.target in compartment_labels
408-
and (edge.source, edge.target) not in back_edges
414+
and edge_pair not in back_edges
409415
):
410416
in_degree[edge.target] += 1
411417
return in_degree
@@ -464,20 +470,27 @@ def _propagate_layers(
464470
Edges to skip.
465471
"""
466472
processed: set[str] = set()
473+
seen_edges: set[tuple[str, str]] = set()
467474
while queue:
468475
node = queue.popleft()
469476
if node in processed:
470477
continue
471478
processed.add(node)
472479

473480
for neighbor in sorted(graph.successors(node)):
474-
if neighbor not in compartment_labels or (node, neighbor) in back_edges:
481+
edge_pair = (node, neighbor)
482+
if (
483+
neighbor not in compartment_labels
484+
or edge_pair in back_edges
485+
or edge_pair in seen_edges
486+
):
475487
continue
488+
seen_edges.add(edge_pair)
476489
new_layer = layers[node] + 1
477490
if neighbor not in layers or new_layer > layers[neighbor]:
478491
layers[neighbor] = new_layer
479-
in_degree[neighbor] -= 1
480-
if in_degree[neighbor] <= 0:
492+
in_degree[neighbor] = max(0, in_degree[neighbor] - 1)
493+
if in_degree[neighbor] == 0:
481494
queue.append(neighbor)
482495

483496

@@ -552,7 +565,22 @@ def sort_key(label: str) -> tuple[int, int, str]:
552565
else:
553566
barycenters[label] = 0.0
554567

555-
return sorted(nodes_in_layer, key=lambda n: (barycenters[n], n))
568+
def barycenter_sort_key(label: str) -> tuple[float, str]:
569+
"""Return sort key based on barycenter and label.
570+
571+
Parameters
572+
----------
573+
label : str
574+
Node label.
575+
576+
Returns
577+
-------
578+
tuple[float, str]
579+
Tuple of barycenter value and label for sorting.
580+
"""
581+
return barycenters[label], label
582+
583+
return sorted(nodes_in_layer, key=barycenter_sort_key)
556584

557585

558586
def _node_sort_index(label: str) -> float:
@@ -566,9 +594,10 @@ def _node_sort_index(label: str) -> float:
566594
Returns
567595
-------
568596
float
569-
The sort index based on label hash.
597+
The sort index based on deterministic label hash.
570598
"""
571-
return float(hash(label) % 1000) / 1000.0
599+
digest = hashlib.md5(label.encode(), usedforsecurity=False).digest()
600+
return int.from_bytes(digest[:4], "big") / 4294967296.0
572601

573602

574603
def _find_back_edges(graph: KineticGraph, compartment_labels: set[str]) -> set[tuple[str, str]]:

tests/inspect/kinetic_scheme/test_k_matrix_parser.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
from __future__ import annotations
44

5+
from copy import deepcopy
6+
57
import pytest
68
from glotaran.testing.simulated_data.parallel_spectral_decay import SCHEME as SCHEME_PAR
79
from glotaran.testing.simulated_data.sequential_spectral_decay import SCHEME as SCHEME_SEQ
@@ -175,3 +177,14 @@ def test_exclude_megacomplexes(self) -> None:
175177
exclude_megacomplexes={"megacomplex_sequential_decay"},
176178
)
177179
assert len(transitions) == 0
180+
181+
def test_missing_dataset_megacomplex_raises(self) -> None:
182+
"""Dataset referencing undefined megacomplex raises ValueError."""
183+
model = deepcopy(SCHEME_SEQ.model)
184+
model.dataset["dataset_1"].megacomplex = ["missing_mc"]
185+
186+
with pytest.raises(
187+
ValueError,
188+
match="referenced by dataset 'dataset_1' not found in model",
189+
):
190+
extract_dataset_transitions("dataset_1", model, SCHEME_SEQ.parameters)

tests/inspect/kinetic_scheme/test_layout.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,15 @@
22

33
from __future__ import annotations
44

5+
import hashlib
6+
57
import pytest
68

79
from pyglotaran_extras.inspect.kinetic_scheme._k_matrix_parser import Transition
810
from pyglotaran_extras.inspect.kinetic_scheme._kinetic_graph import KineticGraph
911
from pyglotaran_extras.inspect.kinetic_scheme._layout import LayoutAlgorithm
1012
from pyglotaran_extras.inspect.kinetic_scheme._layout import _find_connected_components
13+
from pyglotaran_extras.inspect.kinetic_scheme._layout import _node_sort_index
1114
from pyglotaran_extras.inspect.kinetic_scheme._layout import compute_layout
1215

1316

@@ -115,6 +118,15 @@ def test_deterministic_output(self) -> None:
115118
for label in pos1:
116119
assert pos1[label] == pos2[label]
117120

121+
def test_node_sort_index_is_deterministic(self) -> None:
122+
"""Node sort index should use deterministic hashing."""
123+
label = "species_2"
124+
digest = hashlib.md5(label.encode(), usedforsecurity=False).digest()
125+
expected = int.from_bytes(digest[:4], "big") / 4294967296.0
126+
actual = _node_sort_index(label)
127+
assert actual == expected
128+
assert 0.0 <= actual < 1.0
129+
118130
def test_parallel_nodes_side_by_side(self) -> None:
119131
"""Parallel decay nodes (all isolated) should be on the same row."""
120132
graph = _make_parallel_graph()

0 commit comments

Comments
 (0)