Skip to main content

Deep Jointly-Informed Neural Networks

DJINN: Deep jointly-informed neural networks

Tests Formatting codecov Docs License: BSD Python

Fork notice: This repository is a fork and continuation of LLNL's DJINN project (Deep Jointly-Informed Neural Networks) originally developed by Kelli D. Humbird (humbird1@llnl.gov). The original project is available at https://github.com/LLNL/djinn and is distributed under the license found in LICENSE. This fork is maintained by Ben Whewell (ben.whewell@pm.me) — https://github.com/bwhewe-13/DJINN

DJINN is an easy-to-use algorithm for training deep neural networks on supervised regression tasks. For additional information, refer to the paper "Deep neural network initialization with decision trees", cited below.

Getting Started

Original DJINN required TensorFlow. This fork is implemented with PyTorch.

Requirements:

  • Python 3.10+
  • PyTorch
  • scikit-learn

Install from source:

git clone https://github.com/bwhewe-13/DJINN.git
cd DJINN
python -m pip install --upgrade pip
python -m pip install .

Try it out using the examples in examples:

cd examples
python djinn_regression.py
python djinn_classification.py
python djinn_multiout.py

Notes:

  • The scikit-learn version used when training a DJINN model should match the version used when loading/evaluating that saved model.

  • Some example workflows may require matplotlib:

    python -m pip install matplotlib
    

Development

Set up a local development environment:

python -m pip install --upgrade pip
python -m pip install -e .[dev]

Run quality checks and tests:

black --check djinn tests examples
isort --check-only djinn tests examples
flake8
pytest

Enable pre-commit hooks (optional, recommended):

pre-commit install
pre-commit run --all-files

Build docs locally:

python -m pip install sphinx
cd docs
make html

Documentation

To view the DJINN documentation:

cd docs
make html

Open docs/_build/html/index.html in a browser

Source Repo Verification

These tests verify that this PyTorch fork produces results consistent with the original TensorFlow DJINN implementation.

The verification suite has two layers:

  • compare_results.py — an exploratory reporting script that prints a human-readable comparison of two JSON result files. Useful for investigating differences interactively.
  • tests/test_tf_comparison.py — pytest tests that formally gate the comparison. These use asymmetric thresholds: they only fail when PT is worse than TF, not when PT is better. This avoids false failures caused by TF's known convergence instability on certain seeds.

One-time setup

Run setup_envs.sh from the repo root to create both virtual environments:

bash setup_envs.sh

This clones the TF repo into repos/DJINN-tf/ and installs both environments:

  • venvs/tf-djinn/ — TensorFlow implementation
  • venvs/pt-djinn/ — PyTorch implementation (installed from the current repo)

Step 1: Unit tests (run in both envs)

These tests check API compatibility, output shapes, determinism, and save/load correctness. Run them independently in each environment:

# TensorFlow env — shared contract only
source venvs/tf-djinn/bin/activate
pytest tests/test_unit_shared.py -v
deactivate

# PyTorch env — shared contract + PT-specific behaviour
source venvs/pt-djinn/bin/activate
pytest tests/test_unit_shared.py tests/test_unit.py -v
deactivate

Any test that fails in one environment but passes in the other reveals a behavioral divergence between the two implementations.

Step 2: Collect benchmark results

Run the benchmark script once in each environment. This trains DJINN across 20 random seeds and saves the metrics to JSON. The TF results are committed to the repo as a baseline; only the PT results need to be regenerated on each comparison run.

TF baseline (results_tf.json)

results_tf.json is committed to the repo root and should be treated as stable. Only regenerate it if the TF implementation, datasets, or collection methodology change. It was generated with:

  • TensorFlow version: 2.21.0
  • Command: python run_and_collect.py --impl tf --out results_tf.json --ntrees 3 --epochs 100 --seeds 20

To regenerate:

source venvs/tf-djinn/bin/activate
python run_and_collect.py --impl tf --out results_tf.json --ntrees 3 --epochs 100 --seeds 20
git add results_tf.json
git commit -m "Regenerate TF baseline (TF 2.21.0, ntrees=3, epochs=100, seeds=20)"
deactivate
# TensorFlow env — generate committed baseline (one-time)
source venvs/tf-djinn/bin/activate
python run_and_collect.py --impl tf --out results_tf.json --ntrees 3 --epochs 100
deactivate

# PyTorch env — regenerate on each comparison run
source venvs/pt-djinn/bin/activate
python run_and_collect.py --impl pt --out results_pt.json --ntrees 3 --epochs 100
deactivate

Step 3: Compare results

Exploratory report (human-readable, no pass/fail gates):

python compare_results.py --tf results_tf.json --pt results_pt.json

# Optional: generate distribution plots (requires matplotlib)
python compare_results.py --tf results_tf.json --pt results_pt.json --plot

Formal pytest comparison (used in CI):

source venvs/pt-djinn/bin/activate
pytest tests/test_tf_comparison.py -m comparison -v

Interpreting compare_results.py output

compare_results.py is an exploratory tool. Its PASS/WARN/FAIL labels use two-sided statistical tests and are intended to guide investigation, not to formally gate correctness. See test_tf_comparison.py for the authoritative pass/fail criteria used in CI.

Color / Status Meaning
Green PASS Distributions are not significantly different (KS + Mann-Whitney p > 0.05)
Yellow WARN Marginal difference — worth investigating but not necessarily a bug
Red FAIL Statistically significant difference or metric below performance floor

A FAIL in compare_results.py does not necessarily mean the PT implementation is wrong. PT often scores better than TF (higher R**2, lower MSE, lower variance across seeds), which also triggers a FAIL under two-sided tests. Use the pytest suite to determine whether a difference is a real regression.

Acceptance thresholds

These are the criteria used in tests/test_tf_comparison.py. All checks are asymmetric: PT is only required to be at least as good as TF, not identical to it.

Check Threshold
Network architecture Exact match
Prediction shape Exact match
Prediction dtype Must be float (float32 vs float64 acceptable)
Same-seed determinism rtol=1e-4
Save/load round-trip rtol=1e-5
PT median R² not below TF PT median ≥ TF median − 0.05
PT R² not stochastically worse One-sided Mann-Whitney p > 0.05
PT R² variance PT std ≤ TF std × 2
PT multi-output median R² PT median ≥ TF median − 0.02
PT multi-output variance PT std ≤ TF std × 3
BMA uncertainty > 0 Required for all seeds
BMA uncertainty ratio PT/TF Between 0.05× and 10×
BMA output shape Exact match between implementations
Batch size Exact match (data-driven, should be identical)
Learning rate In range [1e-6, 1.0]

Notes

  • The two implementations will never produce identical outputs — TF and PyTorch have different RNGs and optimizer defaults.
  • The goal is that PT is at least as good as TF, not that they are numerically identical.
  • PT consistently converges more reliably than TF (lower R² variance across seeds), particularly on multi-output tasks.
  • Use at least 100 epochs for meaningful comparison results.
  • results_tf.json is committed to the repo root and must not be deleted or regenerated casually — it was produced with TensorFlow 2.21.0 using --ntrees 3 --epochs 100 --seeds 20. Regenerate only when the TF implementation or datasets change, and update the commit message with the TF version and command used.
  • results_pt.json should never be committed — add it to .gitignore.

Source Repo

DJINN is available at https://github.com/LLNL/DJINN

Citing DJINN

If you use DJINN in your research, please cite the following paper:

K. D. Humbird, J. L. Peterson and R. G. Mcclarren, "Deep Neural Network Initialization With Decision Trees," in IEEE Transactions on Neural Networks and Learning Systems, vol. 30, no. 5, pp. 1286-1295, May 2019. doi: 10.1109/TNNLS.2018.2869694, URL: http://ieeexplore.ieee.org/stamp/stamp.jsp?tp=&arnumber=8478232&isnumber=8695188

Release

Copyright (c) 2018, Lawrence Livermore National Security, LLC.

Produced at the Lawrence Livermore National Laboratory

Written by K. Humbird (humbird1@llnl.gov), L. Peterson (peterson76@llnl.gov).

LLNL-CODE-754815 OCEC-18-117

All rights reserved.

Unlimited Open Source- BSD Distribution.

For release details and restrictions, please read the RELEASE, LICENSE, and NOTICE files, linked below:

Metadata

Release files for djinnml 1.1.1

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

Source distribution (sdist)

Source distribution for djinnml 1.1.1
File Size Uploaded
djinnml-1.1.1.tar.gz 45.1 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for djinnml 1.1.1
File Interpreter ABI Platform
djinnml-1.1.1-py3-none-any.whl Python 3 none any Details

Total release size: 75.2 kB

Release files / djinnml-1.1.1.tar.gz

Download URL djinnml-1.1.1.tar.gz
Size 45.1 kB
Tags Source
SHA-256 checksum
How to use checksums
9bb5a650ef3947dbb3a2cab6e619f2c594f30295d630ba6bf6e4f771312a499f
BLAKE2b-256 checksum
How to use checksums
ce8929eb82a7496fd3872cfe09545ea58e2b864849c5dcc419e394b2d58501e4
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 5, 2026.

Transparency log

Release files / djinnml-1.1.1-py3-none-any.whl

Download URL djinnml-1.1.1-py3-none-any.whl
Size 30.1 kB
Tags Python 3
SHA-256 checksum
How to use checksums
be054c34236f7c4c6dd7491728014966fa2e2e19a1d2d1a3dd2c991fede957cf
BLAKE2b-256 checksum
How to use checksums
6b4e8b0526fca394b6e8c9bfa4dfe8248c665b22f78807198e4a302d4b5dc8fd
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 5, 2026.

Transparency log

Release history Release notifications | RSS feed

1.1.2

2 release files

This release

1.1.1 This release

2 release files

1.1.0

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