Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion circulax/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,13 @@
from circulax._version import __version__
from circulax.circuit import Circuit, compile_circuit
from circulax.compiler import compile_netlist
from circulax.netlist import build_net_map, build_net_map_kfnetlist, netlist, sax_to_kfnetlist
from circulax.netlist import (
build_net_map,
build_net_map_kfnetlist,
flatten_recursive_netlist,
netlist,
sax_to_kfnetlist,
)
from circulax.netlist import circulaxNetlist as Netlist
from circulax.s_transforms import fdomain_component, sax_component
from circulax.solvers import analyze_circuit, setup_ac_sweep, setup_harmonic_balance, setup_transient
Expand Down
94 changes: 83 additions & 11 deletions circulax/circuit.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,17 +24,11 @@ def _resolve_osdi_param_col(group: Any, param_key: str) -> int:
model = get_model(group.model_id)
name_to_col = {n.lower(): i for i, n in enumerate(model.param_names)}
except (ImportError, Exception) as exc:
msg = (
f"Cannot resolve OSDI parameter '{param_key}': "
f"bosdi registry lookup failed ({exc!r})."
)
msg = f"Cannot resolve OSDI parameter '{param_key}': bosdi registry lookup failed ({exc!r})."
raise ValueError(msg) from exc
col = name_to_col.get(param_key.lower())
if col is None:
msg = (
f"Parameter '{param_key}' not found in OSDI model "
f"(available: {sorted(name_to_col)})."
)
msg = f"Parameter '{param_key}' not found in OSDI model (available: {sorted(name_to_col)})."
raise ValueError(msg)
return col

Expand Down Expand Up @@ -82,6 +76,8 @@ def __init__(
rtol: float = 1e-6,
atol: float = 1e-6,
max_steps: int = 100,
_source_netlist: dict | None = None,
_source_models: dict | None = None,
) -> None:
self.solver = solver
self.groups = groups
Expand All @@ -90,6 +86,25 @@ def __init__(
self.rtol = rtol
self.atol = atol
self.max_steps = max_steps
self._source_netlist = _source_netlist
self._source_models = _source_models

@property
def ports(self) -> tuple[str, ...]:
"""External port names declared in the source netlist."""
if self._source_netlist is None:
return ()
return tuple(self._source_netlist.get("ports", {}).keys())

@property
def source_netlist(self) -> dict | None:
"""The original netlist used to compile this circuit, if available."""
return self._source_netlist

@property
def source_models(self) -> dict | None:
"""The leaf models used to compile this circuit, if available."""
return self._source_models

def _n(self) -> int:
return self.sys_size * (2 if self.solver.is_complex else 1)
Expand Down Expand Up @@ -453,6 +468,36 @@ def with_groups(self, groups: dict) -> Circuit:
)


def _embed_circuit_subcircuits(
net_dict: dict | kfnl.Netlist,
models_map: dict,
circuit_models: dict[str, Circuit],
) -> dict:
"""Build a RecursiveNetlist from Circuit objects in *models_map* (mutates *models_map*)."""
from circulax.netlist import _is_recursive_netlist

if isinstance(net_dict, kfnl.Netlist):
net_dict = net_dict.to_dict()
recnet: dict[str, dict] = {}
if isinstance(net_dict, dict) and _is_recursive_netlist(net_dict):
recnet.update(net_dict)
else:
recnet["top"] = net_dict # type: ignore[assignment]
for name, circ in circuit_models.items():
if circ.source_netlist is None:
msg = f"Circuit '{name}' has no stored source netlist and cannot be used as a subcircuit."
raise ValueError(msg)
recnet[name] = circ.source_netlist
for mk, mv in (circ.source_models or {}).items():
existing = models_map.get(mk)
if existing is not None and existing is not mv:
msg = f"Model name conflict: '{mk}' maps to different objects in parent and subcircuit '{name}'."
raise ValueError(msg)
models_map[mk] = mv
del models_map[name]
return recnet


def compile_circuit(
net_dict: dict | kfnl.Netlist,
models_map: dict,
Expand All @@ -466,11 +511,19 @@ def compile_circuit(
) -> Circuit:
"""Compile a netlist into a callable :class:`Circuit`.

Accepts either a ``kfnetlist.Netlist`` or a SAX-format dict.
Accepts a ``kfnetlist.Netlist``, a SAX-format dict, or a
``RecursiveNetlist`` (``dict[str, Netlist]``). When a recursive netlist
is given, subcircuit instances are flattened before compilation.

A compiled :class:`Circuit` may also appear as a value in *models_map*;
its stored source netlist is inlined as a subcircuit automatically.

Args:
net_dict: Netlist (kfnetlist.Netlist or SAX-format dict).
models_map: Mapping from component type name strings to component classes.
net_dict: Netlist (kfnetlist.Netlist, SAX-format dict, or
RecursiveNetlist).
models_map: Mapping from component type name strings to component
classes, SAX model functions, or compiled :class:`Circuit`
objects.
backend: Linear solver backend (``"default"``, ``"dense"``, ``"klu"`` etc.).
is_complex: If ``True``, treat the circuit as complex-valued (photonic).
If ``"auto"`` (default), infer this from component outputs.
Expand All @@ -484,8 +537,25 @@ def compile_circuit(

"""
from circulax.compiler import compile_netlist
from circulax.netlist import _is_recursive_netlist, flatten_recursive_netlist
from circulax.solvers.linear import analyze_circuit

models_map = dict(models_map)
source_netlist: dict | None = None
source_models: dict | None = None

circuit_models = {k: v for k, v in models_map.items() if isinstance(v, Circuit)}
if circuit_models:
net_dict = _embed_circuit_subcircuits(net_dict, models_map, circuit_models)

if isinstance(net_dict, dict) and _is_recursive_netlist(net_dict):
source_netlist = net_dict.get(next(iter(net_dict)))
source_models = {k: v for k, v in models_map.items() if not isinstance(v, Circuit)}
net_dict = flatten_recursive_netlist(net_dict)
elif isinstance(net_dict, dict):
source_netlist = net_dict
source_models = {k: v for k, v in models_map.items() if not isinstance(v, Circuit)}

groups, sys_size, port_map = compile_netlist(net_dict, models_map)
if is_complex == "auto":
is_complex = _infer_is_complex(groups)
Expand All @@ -501,6 +571,8 @@ def compile_circuit(
rtol=rtol,
atol=atol,
max_steps=max_steps,
_source_netlist=source_netlist,
_source_models=source_models,
)


Expand Down
127 changes: 127 additions & 0 deletions circulax/netlist.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,133 @@ def _net_member_ref(port_str: str) -> kfnl.PortRef | kfnl.NetlistPort:
return nl, settings_override


# ---------------------------------------------------------------------------
# Recursive netlist flattening
# ---------------------------------------------------------------------------


def _is_recursive_netlist(net_dict: dict) -> bool:
"""Return True if *net_dict* looks like a RecursiveNetlist (dict-of-Netlists)."""
if "instances" in net_dict:
return False
return any(isinstance(v, dict) and "instances" in v for v in net_dict.values())


def flatten_recursive_netlist(
recnet: dict[str, dict],
sep: str = "~",
) -> dict:
"""Flatten a ``RecursiveNetlist`` into a single SAX-format netlist.

A ``RecursiveNetlist`` is a ``dict[str, Netlist]`` where the first key is
the top-level circuit and remaining keys define subcircuits. An instance
whose ``component`` string matches a key in the dict is treated as a
subcircuit and inlined with prefixed instance names.

Handles circulax connection extensions (tuple targets, ``nets`` lists)
that SAX's own ``flatten_netlist`` does not support.

Args:
recnet: Mapping from circuit name to SAX-format netlist dict.
sep: Separator for hierarchical instance names (default ``"~"``).

Returns:
A flat SAX-format netlist dict.

"""
import copy

top_name = next(iter(recnet))
flat = copy.deepcopy(recnet[top_name])
_flatten_into(recnet, flat, sep)
return flat


def _rewrite_ref(ref: str, inst_name: str, port_map: dict[str, str]) -> str:
"""Rewrite a port reference if it targets *inst_name*."""
if "," not in ref:
return ref
inst, port = ref.split(",", 1)
if inst == inst_name:
mapped = port_map.get(port)
if mapped is not None:
return mapped
return ref


def _rewrite_connection_value(
val: str | tuple | list,
inst_name: str,
port_map: dict[str, str],
) -> str | tuple:
"""Rewrite the value side of a connection entry."""
if isinstance(val, str):
return _rewrite_ref(val, inst_name, port_map)
return tuple(_rewrite_ref(v, inst_name, port_map) for v in val)


def _flatten_into(recnet: dict[str, dict], net: dict, sep: str) -> None:
"""Inline all subcircuit instances in *net* (mutates in place)."""
import copy

changed = True
while changed:
changed = False
for inst_name in list(net.get("instances", {})):
comp = net["instances"][inst_name].get("component", "")
if comp not in recnet:
continue
changed = True
child = copy.deepcopy(recnet[comp])
_flatten_into(recnet, child, sep)
_inline_subcircuit(net, inst_name, child, sep)


def _inline_subcircuit(net: dict, inst_name: str, child: dict, sep: str) -> None:
"""Inline a single flattened subcircuit into *net* (mutates in place)."""
del net["instances"][inst_name]

port_map: dict[str, str] = {ext: _prefix_ref(ref, inst_name, sep) for ext, ref in child.get("ports", {}).items()}

for child_inst, child_data in child.get("instances", {}).items():
if child_inst == "GND" or child_data.get("component") == "ground":
net["instances"].setdefault("GND", child_data)
else:
net["instances"][f"{inst_name}{sep}{child_inst}"] = child_data

connections = net.setdefault("connections", {})
for src, tgt in child.get("connections", {}).items():
new_src = _prefix_ref(src, inst_name, sep)
if isinstance(tgt, str):
new_tgt: str | tuple = _prefix_ref(tgt, inst_name, sep)
else:
new_tgt = tuple(_prefix_ref(t, inst_name, sep) for t in tgt)
connections[new_src] = new_tgt

for net_entry in child.get("nets", []):
p1 = _prefix_ref(net_entry["p1"], inst_name, sep)
p2 = _prefix_ref(net_entry["p2"], inst_name, sep)
net.setdefault("nets", []).append({"p1": p1, "p2": p2})

net["connections"] = {
_rewrite_ref(src, inst_name, port_map): _rewrite_connection_value(tgt, inst_name, port_map)
for src, tgt in connections.items()
}

if "ports" in net:
net["ports"] = {pname: _rewrite_ref(ptgt, inst_name, port_map) for pname, ptgt in net["ports"].items()}


def _prefix_ref(ref: str, inst_name: str, sep: str) -> str:
"""Prefix an ``"instance,port"`` reference, skipping GND."""
if "," not in ref:
return ref
inst, port = ref.split(",", 1)
if inst == "GND":
return ref
return f"{inst_name}{sep}{inst},{port}"


# ---------------------------------------------------------------------------
# Legacy SAX build_net_map (kept for backward compat / draw_circuit_graph)
# ---------------------------------------------------------------------------
Expand Down
Loading
Loading