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
60 changes: 60 additions & 0 deletions .github/workflows/docs.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
name: Docs

on:
push:
branches: [main]
workflow_dispatch:

permissions:
contents: read
pages: write
id-token: write

concurrency:
group: pages
cancel-in-progress: false

jobs:
build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4

- uses: actions/setup-python@v5
with:
python-version: '3.11'
cache: pip
cache-dependency-path: docs/requirements.txt

- name: Install pandoc
run: sudo apt-get update && sudo apt-get install -y pandoc

- name: Install Python deps
run: pip install -r docs/requirements.txt

- name: Generate tutorial pages from .lhs
run: |
mkdir -p docs/mkdocs/docs/examples
for f in QueueModel OptionPricing GeneTranscription; do
pandoc --from=markdown+lhs --to=commonmark "tutorials/${f}.lhs" \
-o "docs/mkdocs/docs/examples/${f}.md"
done

- name: Build mkdocs site
run: cd docs/mkdocs && mkdocs build

- uses: actions/configure-pages@v5

- uses: actions/upload-pages-artifact@v3
with:
path: docs/mkdocs/site

deploy:
needs: build
runs-on: ubuntu-latest
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
steps:
- id: deployment
uses: actions/deploy-pages@v4
4 changes: 4 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,10 @@ stack.yaml.lock
# Local dev config (kept locally, not shipped)
.hlint.yaml

# mkdocs build outputs
docs/mkdocs/site/
docs/mkdocs/docs/examples/*.md

# Python
docs/.venv/
__pycache__/
Expand Down
25 changes: 13 additions & 12 deletions README.md
Original file line number Diff line number Diff line change
@@ -1,16 +1,19 @@
<p align="center">
<img src="docs/assets/logo.png#gh-light-mode-only" alt="GradInf Logo" height="60">
<img src="docs/assets/logo-white.png#gh-dark-mode-only" alt="GradInf Logo" height="60">
</p>
<div align="center">
<img src="docs/assets/logo.png#gh-light-mode-only" alt="GradInf Logo" height="100">
<img src="docs/assets/logo-white.png#gh-dark-mode-only" alt="GradInf Logo" height="100">
</div>


# GradInf

This repository provides the Haskell implementation
of *gradient inference*, a new approach to gradient
GradInf is a research package that provides the Haskell implementation
of **gradient inference**, a new approach to gradient
estimation described in
[this PLDI'26 paper](https://dl.acm.org/doi/abs/10.1145/3808321).

**Polished documentation and tutorials will be added shortly!**
[![Build Status](https://github.com/probsys/grad-inf/actions/workflows/ci.yml/badge.svg?branch=main)](https://github.com/probsys/grad-inf/actions/workflows/ci.yml)
[![](https://img.shields.io/badge/docs-main-blue.svg)](https://probsys.github.io/grad-inf/)
[![DOI](https://img.shields.io/badge/article-DOI%3A10.1145%2F3808321-B31B1B)](https://dl.acm.org/doi/abs/10.1145/3808321)

## Overview

Expand Down Expand Up @@ -71,13 +74,11 @@ print gradientEstimate

Here, `stratifiedImportanceResamplingInferenceAlg 1` is an *inference strategy*
which enables sound and efficient gradient estimation.
See [Documentation](#documentation) for tutorials and details on the API.

[1] *Monte Carlo Gradient Estimation in Machine Learning*, Shakir Mohamed, Mihaela Rosca, Michael Figurnov, and Andriy Mnih. *Journal of Machine Learning Research* 21, 132 (2020),

## Documentation
To learn more, check out the [documentation](https://probsys.github.io/grad-inf).
Additional tutorials and API details will be added shortly!

Documentation will be added shortly!
[1] *Monte Carlo Gradient Estimation in Machine Learning*, Shakir Mohamed, Mihaela Rosca, Michael Figurnov, and Andriy Mnih. *Journal of Machine Learning Research* 21, 132 (2020),

## Reproducible Artifact

Expand Down
26 changes: 26 additions & 0 deletions docs/Makefile
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
SHELL := /bin/bash
ROOT := ..
TUTORIALS := $(ROOT)/tutorials
MKDOCS_DOCS := mkdocs/docs
EXAMPLES_DST := $(MKDOCS_DOCS)/examples
VENV := .venv

TUTORIAL_NAMES := QueueModel OptionPricing GeneTranscription
TUTORIAL_OUTPUTS := $(addprefix $(EXAMPLES_DST)/, $(addsuffix .md, $(TUTORIAL_NAMES)))

.PHONY: docs serve clean

docs: $(TUTORIAL_OUTPUTS)
. $(VENV)/bin/activate && cd mkdocs && mkdocs build

serve: $(TUTORIAL_OUTPUTS)
. $(VENV)/bin/activate && cd mkdocs && mkdocs serve

$(EXAMPLES_DST)/%.md: $(TUTORIALS)/%.lhs | $(EXAMPLES_DST)
pandoc --from=markdown+lhs --to=commonmark $< -o $@

$(EXAMPLES_DST):
mkdir -p $@

clean:
rm -rf mkdocs/site $(EXAMPLES_DST)/*.md
Binary file modified docs/assets/logo-white.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Binary file modified docs/assets/logo.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
1 change: 1 addition & 0 deletions docs/mkdocs/docs/images/logo-white.png
1 change: 1 addition & 0 deletions docs/mkdocs/docs/images/logo.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
1 change: 1 addition & 0 deletions docs/mkdocs/docs/images/workflow-core.png
97 changes: 97 additions & 0 deletions docs/mkdocs/docs/index.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
<p align="center">
<img src="images/logo.png" alt="GradInf Logo" height="20">
</p>

# GradInf: Gradient Estimation as Probabilistic Inference

GradInf is a research package that provides the Haskell implementation
of *gradient inference*, a new approach to gradient
estimation described in
[this PLDI'26 paper](https://dl.acm.org/doi/abs/10.1145/3808321).

## Overview

<p align="center">
<img src="images/workflow-core.png" alt="Core GradInf Workflow" width="800">
</p>

GradInf automatically synthesizes gradient estimators
(i.e., estimators of gradients of expected values [1]) for
probabilistic programs. Users define a probabilistic
program using a mix of ordinary Haskell and library-provided
primitives. They then call `gradInfAD`, specifying their desired
probabilistic inference strategy. The result is an unbiased estimate of
the gradient of the expectation
of the original program with respect to its parameter.

### Starter Example

We can write a simple queuing model:
```haskell
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE RebindableSyntax #-}

import Prelude hiding (flip)
import Numeric.GradInf.Primitives.DeterministicPrimitives
import Numeric.GradInf.Primitives.FlipCRN
import Numeric.GradInf.Primitives.IterateP

queueKernel :: forall m d i b mat.
(DeterministicPrimitives d i b mat, FlipCRN m d b)
=> d -> i -> m i
queueKernel theta x = do
let p = theta / (theta + if (isGreater x 25) :: b then 25 else toDouble x)
b :: b <- flipCRN p
let x' = if b then x + 1 else x - 1
return x'

queueModel :: forall m d i b mat.
(DeterministicPrimitives d i b mat, FlipCRN m d b, IterateP m i)
=> Int -> d -> m d
queueModel n theta = do
x <- iterateP (queueKernel theta) 0 !! n
return (toDouble x)
```
The annotated probability distributions (e.g. `flipCRN`) specify
*factorized coupling strategies*, which GradInf uses to form
lower variance gradient estimators.
We can now differentiate the program using the GradInf high-level API:
```haskell
import Data.Functor.Identity
import Numeric.GradInf

let thetaToDifferentiateAt = Identity 15.0
let n = 50
gradientEstimate <- sampler $ gradInfAD (queueModel n . runIdentity) (stratifiedImportanceResamplingInferenceAlg 1) thetaToDifferentiateAt
print gradientEstimate
```

Here, `stratifiedImportanceResamplingInferenceAlg 1` is an *inference strategy*
which enables sound and efficient gradient estimation.
See the tutorials in this documentation site for more details on the API.

[1] *Monte Carlo Gradient Estimation in Machine Learning*, Shakir Mohamed, Mihaela Rosca, Michael Figurnov, and Andriy Mnih. *Journal of Machine Learning Research* 21, 132 (2020),

## Index

For more usage examples, check out the tutorial pages on applying GradInf to [queuing models](examples/QueueModel.md), [financial models](examples/OptionPricing.md), and [biochemical models](examples/GeneTranscription.md).
*API documentation and additional tutorials will be added soon!*

## Citation

```bibtex
@article{arya2026gradinf,
title = {GradInf: Gradient Estimation as Probabilistic Inference},
author = {Arya, Gaurav and Huot, Mathieu and Schauer, Moritz and Lew, Alexander K. and Saad, Feras A.},
journal = {Proceedings of the ACM on Programming Languages},
volume = {10},
number = {PLDI},
articleno = {243},
month = jun,
pages = {1864--1890},
year = {2026},
publisher = {Association for Computing Machinery},
address = {New York, NY, USA},
doi = {10.1145/3808321},
}
```
22 changes: 22 additions & 0 deletions docs/mkdocs/docs/overrides.css
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
p {
margin: 0 0 18px;
}

.md-typeset h1, .md-typeset h2, .md-typeset h3, .md-typeset h4, .md-typeset h5, .md-typeset h6 {
margin-bottom: 18px;
color: black;
font-weight: 500;
}

.md-typeset .output {
color: black;
background: white;
}

/* Force the inline index-page logo to its `height` attribute. Material's
`.md-typeset img { height: auto }` otherwise overrides it. */
.md-typeset img[alt="GradInf Logo"] {
height: 80px !important;
width: auto !important;
}

2 changes: 2 additions & 0 deletions docs/mkdocs/docs/overrides.js
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
hljs.configure({languages:[]});
hljs.initHighlightingOnLoad();
37 changes: 37 additions & 0 deletions docs/mkdocs/mkdocs.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
site_name: GradInf

theme:
name: 'material'
home: index
toc_depth: 2
features:
- content.code.copy
- navigation.expand
logo: images/logo-white.png
palette:
primary: black

nav:
- Home: index.md
- Examples:
- Differentiating a Queue Model: examples/QueueModel.md
- Estimating Greeks of Option Pricing Model: examples/OptionPricing.md
- Parameter Inference of Gene Transcription Model: examples/GeneTranscription.md

markdown_extensions:
- pymdownx.arithmatex:
generic: true
- pymdownx.highlight:
anchor_linenums: true
line_spans: __span
pygments_lang_class: true
- pymdownx.inlinehilite
- pymdownx.snippets
- pymdownx.superfences

extra_javascript:
- https://unpkg.com/mathjax@3/es5/tex-mml-chtml.js
- overrides.js

extra_css:
- overrides.css
3 changes: 3 additions & 0 deletions docs/requirements.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
mkdocs
mkdocs-material
pymdown-extensions
Loading
Loading