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_dimspace 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)
| File | Size | Uploaded | |
|---|---|---|---|
| c3po_neuro-0.1.0.tar.gz | 17.1 MB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| 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 logRelease 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