Skip to content

between_types

champalimaud.between_types

Synapse totals between cell types.

synapses_between_types(connections, types)

Synapses from each type to each type, for every ordered pair.

Parameters:

Name Type Description Default
connections DataFrame

Has columns pre_type, post_type, and weight. A row where either type is null is dropped.

required
types sequence of str

The types to tabulate, in the order of the result.

required

Returns:

Type Description
DataFrame

Columns pre_type, post_type, and synapses, one row for every ordered pair of types, the diagonal included, 0 where there is no connection.

Examples:

>>> import polars as pl
>>> connections = pl.DataFrame(
...     {
...         "pre_type": ["A", "A", "B"],
...         "post_type": ["B", "B", "A"],
...         "weight": [10, 2, 7],
...     }
... )
>>> synapses_between_types(connections, ["A", "B"]).rows()
[('A', 'A', 0), ('A', 'B', 12), ('B', 'A', 7), ('B', 'B', 0)]
Source code in champalimaud/between_types.py
def synapses_between_types(
    connections: pl.DataFrame, types: Sequence[str]
) -> pl.DataFrame:
    """Synapses from each type to each type, for every ordered pair.

    Parameters
    ----------
    connections : polars.DataFrame
        Has columns ``pre_type``, ``post_type``, and ``weight``.
        A row where either type is null is dropped.
    types : sequence of str
        The types to tabulate, in the order of the result.

    Returns
    -------
    polars.DataFrame
        Columns ``pre_type``, ``post_type``, and ``synapses``, one row
        for every ordered pair of `types`, the diagonal included, 0
        where there is no connection.

    Examples
    --------
    >>> import polars as pl
    >>> connections = pl.DataFrame(
    ...     {
    ...         "pre_type": ["A", "A", "B"],
    ...         "post_type": ["B", "B", "A"],
    ...         "weight": [10, 2, 7],
    ...     }
    ... )
    >>> synapses_between_types(connections, ["A", "B"]).rows()
    [('A', 'A', 0), ('A', 'B', 12), ('B', 'A', 7), ('B', 'B', 0)]
    """
    every_pair = pl.DataFrame(
        {
            "pre_type": [a for a in types for _ in types],
            "post_type": [b for _ in types for b in types],
        }
    )
    sums = (
        connections.drop_nulls(["pre_type", "post_type"])
        .group_by("pre_type", "post_type")
        .agg(synapses=pl.col("weight").sum())
    )
    return (
        every_pair.join(sums, on=["pre_type", "post_type"], how="left")
        .with_columns(pl.col("synapses").fill_null(0))
        .sort(
            pl.col("pre_type").replace_strict(
                {t: n for n, t in enumerate(types)}
            ),
            pl.col("post_type").replace_strict(
                {t: n for n, t in enumerate(types)}
            ),
        )
    )