Repository navigation
Expand file tree
/
Copy pathpyproject.toml
More file actions
70 lines (63 loc) · 1.77 KB
/
Copy pathpyproject.toml
File metadata and controls
70 lines (63 loc) · 1.77 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
[project]
name = "jopendde"
version = "0.1.0"
description = "JAX/Equinox translation of OpenDDE"
readme = "README.md"
requires-python = ">=3.12"
dependencies = [
"dm-tree>=0.1.10",
"einops>=0.8.0",
"equinox>=0.11.10",
"jaxtyping>=0.2.36",
"numpy>=1.26.3",
]
[project.optional-dependencies]
# Fused NVIDIA cuEquivariance triangle kernels (CUDA 12). Optional: enables the
# triangle-attention kernel via inference.enable_cue_kernels() (a trunk speedup
# that grows with token count). Not imported unless that helper is called.
cue = [
"cuequivariance-jax>=0.10.0",
"cuequivariance-ops-jax-cu12>=0.10.0",
]
[dependency-groups]
jax-cpu = [
"jax[cpu]",
]
jax-cuda = [
"jax[cuda12]",
]
# The PyTorch OpenDDE reference, vendored under ./OpenDDE. jopendde has no
# featurizer of its own and the weights ship as a torch checkpoint, so
# converting a model / featurizing inputs (Predictor.from_checkpoint,
# Predictor.featurize) need `opendde` + torch. torch/torchvision/torchaudio are
# listed here so the CPU-index source below applies to them.
reference = [
"opendde",
"torch",
"torchvision",
"torchaudio",
]
[tool.uv.sources]
opendde = { path = "OpenDDE", editable = true }
# torch is only used on CPU (weight extraction); pull the CPU build so it
# doesn't drag in ~5 GB of bundled CUDA wheels. JAX owns the GPU.
torch = { index = "pytorch-cpu" }
torchvision = { index = "pytorch-cpu" }
torchaudio = { index = "pytorch-cpu" }
[[tool.uv.index]]
name = "pytorch-cpu"
url = "https://download.pytorch.org/whl/cpu"
explicit = true
[tool.uv]
package = true
conflicts = [
[
{ group = "jax-cpu" },
{ group = "jax-cuda" },
],
]
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["src/jopendde"]