Skip to main content

Contrastive Prediction of Point Process Observations (C3PO)

Unsupervised model to infer latent generative processes from point processes occuring in $\mathbb{R}^d$ spcae. Developed with focus for inferring neural states from waveform features of unclustered spike events.

Analysis associated with the publication "Learning predictive latent structure from unclustered neural spike trains" available at the publication repository

Usage

Model declaration

In both intuition and implementation, the c3po model is broken into three modules. Details and derivation of the overall architecture and role of each of these components is described in the documentation:

  • encoder model: Accepts the waveform feature vector $W$ and embeds it into a lower latent_dim space as the embedded mark vector $Z$.
  • context model: Embeds the history of marks $\set{Z_j}_{j=1}^i$ into a context feature vector $C_i$
  • rate-prediction model: Estimates the set of context-dependent parameters $\theta(Z_i,C_j)$ used to calculate the hazard function of a given embedded waveform $H(\theta(Z_j, C_i),\Delta t_i)$. Implemented options for the emission process model are found here.

The complete model is agnostic to the deep learning architecture used to implement each of these modules. For example, the context model can be implemented as any history- coding model including a RNN, a wavenet architecture, or causal transformers. To specify which of the the implemented versions of each module to use, each has a factory constructor which accepts a key for the architecture type and a dictionary or arguments for the module (e.g. number of layers, convolutional filter size, etc.).

A complete c3po model for training is defined by the C3PO class. The initialization of the class requires constructor arguments for each of the modules and definition of the embedded dimension size. An example is given below.

# hyperparams
latent_dim = 10
context_dim = 10
# encoder
encoder_widths = [128, 128, 64]
encoder_args = dict(
    encoder_model="simple",
    widths=encoder_widths,
)
# context
dilations = [1, 2, 4, 8, 16]
kernels = [10, 20, 64, 64, 128]
dilations = dilations * 2
kernels = kernels * 2
context_args = dict(
    context_model="wavenet",
    layer_dilations=dilations,
    layer_kernel_size=kernels,
    expanded_dim=32,
    smoothing=10,
)
# rate model
rate_args = dict(
    rate_model="bilinear",
)
distribution = "poisson"
n_neg_samples = 128
model = C3PO(
    encoder_args,
    context_args,
    rate_args,
    distribution,
    latent_dim,
    context_dim,
    n_neg_samples,
)

Loss functions

Note: implementation is correct and functional but require code cleanupof naming and documentation. For derivation and explanation of loss functions see documentation

NCE (Recommended): The noise contrastive estimation loss is given by: $\mathcal{L} = - \mathbb{E}_i[\log \frac{H_i/H'_i}{\sum_j H_j/H'_j}]$. Example execution of the C3PO method shown below:

#define model
model = C3PO(
    encoder_args,
    context_args,
    rate_args,
    distribution,
    latent_dim,
    context_dim,
    n_neg_samples,
    return_embeddings_in_call=True,
)
run_model = jax.jit(model.apply)

# forward pass to get \theta params and embedded values
pos_params, neg_params, z, c, neg_z = run_model(params, x, delta_t, rand_key)
# loss
loss =  model.contrastive_sequence_loss(
    pos_params,
    neg_params,
    delta_t,
    z[:, 1:],
    neg_z,
)

MLE: Maximum likelihood loss is given by $\mathcal{L}=\mathbb{E}_i[-\log H_i+\log\bar{S}_i]$ and is provided as a method of the C3PO class. Example execution shown below:

# define model
model = C3PO(
    encoder_args,
    context_args,
    rate_args,
    distribution,
    latent_dim,
    context_dim,
    n_neg_samples,
    predicted_sequence_length,
)
run_model = jax.jit(model.apply)
# forward pass to get \theta params
pos_params, neg_params = run_model(params, x, delta_t, rand_key)
# loss
loss = model.loss_generalized_model(
    pos_params, neg_params, delta_t,
)

Similar Methods

  • Contrastive Predictive Coding (CPC)
    • Works for continuous time rather than point processes
  • Temporal Neural Networks (e.g. Hawkes process)
    • designed for categorically labeled events
  • CEBRA
    • designed for sorted spiking data
    • semi-supervised definition of matching states

Metadata

Release files for c3po-neuro 0.1.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for c3po-neuro 0.1.0
File Size Uploaded
c3po_neuro-0.1.0.tar.gz 17.1 MB Details

Built distribution (wheel)

Table of built distributions (wheels) for c3po-neuro 0.1.0
File Interpreter ABI Platform
c3po_neuro-0.1.0-py3-none-any.whl Python 3 none any Details

Total release size: 17.2 MB

Release files / c3po_neuro-0.1.0.tar.gz

Download URL c3po_neuro-0.1.0.tar.gz
Size 17.1 MB
Tags Source
SHA-256 checksum
How to use checksums
3c487fb410ef49960cb08f9aad01a2b120a4056914f88a6b7b4847789f8a4728
BLAKE2b-256 checksum
How to use checksums
7270e5644aaf22bf288fa8d7e36375659d4bb4159bcac14497b8e55ff6c5e2c1
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Oct 6, 2026.

Transparency log

Release files / c3po_neuro-0.1.0-py3-none-any.whl

Download URL c3po_neuro-0.1.0-py3-none-any.whl
Size 72.9 kB
Tags Python 3
SHA-256 checksum
How to use checksums
c542ed4a8699d14c234900494b697b94242ea2b3e6f01860d52f41fc713873c9
BLAKE2b-256 checksum
How to use checksums
a8897fa2823f81c6f5419e22ca0ca126744bace7453037a05b82ce53c1ce46f3
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Oct 6, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.1.0 This release

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page