patrick-kidger / patrick-kidger/optimistix

Scalable optimization for large data

Open
#201 3 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

question
Dominant language
Python
Stars
623
Forks
54
PR merge metrics
No merged PRs in 30d

Description

First a big thanks to all maintainers for this great library. I am trying to optimize a cost function on potentially larger scale input data for cases where classical mini-batch optimizers converge very slowly. For mini-batch optimizers (e.g. optax), this could be easily done using a dataloader (e.g. grain). However, it seems a bit tricky to do in optimistix. Hoping for some insights and/or ideas how to achieve this with optimistix.

Below is an MWE and a list of my considered variants

  1. gpu sharding (needs more physical gpus)
  2. move to cpu (plenty of mem but slow)
  3. alternating solves for different data batches (needs some logic to ensure convergence)
  4. alternating solves for different nn layers (needs some logic to ensure convergence)
  5. dataloader to calculate cost function iteratively (seems ideal but difficult to integrate)
import jax
import jax.numpy as jnp
import optimistix as optx
import jax.random as jr
from jax.example_libraries import stax

jax.config.update("jax_enable_x64", True)

key = jr.PRNGKey(0)
dim = 10
num_samples = 100_000 # some large number

X = jr.normal(key, (num_samples, dim))
y = jr.normal(key, (num_samples, 1))

activation = stax.Sigmoid
layer_sizes = [80,]*3
init_fun, predict_fun = stax.serial(
    stax.serial( *(sum([[stax.Dense(size), activation] for size in layer_sizes],[],)) + [stax.Dense(1)] ) )
_, params = init_fun(key, (X.shape[1],))

def loss(params, args):
    X, y = args # <- ideally an iterator + fori loop over data batches here
    return jnp.mean( (predict_fun(params, X).flatten() - y.flatten())**2 )

bfgs_tol = 1e-12
solver = optx.BFGS(rtol=bfgs_tol, atol=bfgs_tol)
sol = optx.minimise(
    loss,
    solver,
    params,
    max_steps=100_000,
    throw=False,
    args=(X,y) # <- ideally a dataloader here
)

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with the MWE and the optx.minimise call, then read how optimistix passes args into the loss during optimisation. Compare the listed batching, sharding, and alternating-solve ideas against the current API. Done would require a decided design for scalable data loading and convergence behavior, plus a demonstrated implementation path.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning, performance
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.