Exploration-Distillation (ExpDis) checkpoints

Checkpoints for Decoupling Exploration from Optimization in RLVR (Punjwani and Goldblum). Training code: SaifPunjwani/Exploration-Distillation.

The repository holds eight checkpoints trained from Qwen3-1.7B, Qwen3-4B, and Ministral-3-3B-Instruct-2512. Each folder is named after the method label used in the paper.

Models

Base model Method (paper label) Folder
Qwen3-1.7B ExpDis qwen3-1.7b-expdis
Qwen3-1.7B ExpDis (single-round) qwen3-1.7b-expdis-single-round
Qwen3-1.7B DAPO (4× steps) qwen3-1.7b-dapo-4x-steps
Qwen3-1.7B DAPO qwen3-1.7b-dapo
Qwen3-4B ExpDis qwen3-4b-expdis
Qwen3-4B ExpDis (single-round) qwen3-4b-expdis-single-round
Ministral-3-3B-Instruct-2512 ExpDis ministral-3-3b-expdis
Ministral-3-3B-Instruct-2512 ExpDis (single-round) ministral-3-3b-expdis-single-round

ExpDis is the multi-round, multi-explorer configuration. ExpDis (single-round) uses one explorer and one round. DAPO is the correctness-only baseline, and DAPO (4× steps) is the same baseline trained for four times as many steps.

Directory layout

Each folder has two subfolders with the same weights:

  • gpu/: safetensors, loadable with Transformers.
  • tpu/: a flat dictionary of Hugging Face-named parameters for JAX. Most folders store it as flax_model.msgpack. qwen3-1.7b-expdis and qwen3-1.7b-dapo store it as safetensors and include load_params.py.

The repository contains inference weights only.

Transformers inference

pip install "transformers>=5.16.1" "mistral-common>=1.11.7" torch
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

repo_id = "SaifPunjwani/expdis-checkpoints"
subfolder = "qwen3-1.7b-expdis/gpu"

device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if device == "cuda" else torch.float32

tokenizer = AutoTokenizer.from_pretrained(repo_id, subfolder=subfolder)
model = AutoModelForCausalLM.from_pretrained(
    repo_id,
    subfolder=subfolder,
    dtype=dtype,
).to(device)

system = r"Please reason step by step, and put your final answer within \boxed{}."
messages = [{"role": "system", "content": system},
            {"role": "user", "content": "Compute 2+2."}]
inputs = tokenizer.apply_chat_template(
    messages,
    tokenize=True,
    add_generation_prompt=True,
    return_tensors="pt",
    return_dict=True,
    enable_thinking=True,  # Qwen3; omit for Ministral
).to(device)

output = model.generate(**inputs, max_new_tokens=64, do_sample=False)
print(tokenizer.decode(
    output[0, inputs["input_ids"].shape[1]:],
    skip_special_tokens=True,
))

Set subfolder to any gpu/ path in the table above.

JAX/TPU loading

To restore the parameters from a MsgPack export:

from flax.serialization import msgpack_restore
from huggingface_hub import hf_hub_download

path = hf_hub_download(
    "SaifPunjwani/expdis-checkpoints",
    "qwen3-1.7b-expdis-single-round/tpu/flax_model.msgpack",
)
with open(path, "rb") as handle:
    params = msgpack_restore(handle.read())

These are flat parameter dictionaries, not FlaxAutoModelForCausalLM directories. For the Transformers generation API, use the gpu/ subfolder.

jax_runtime/ is a small inference loader for all eight tpu/ exports. It reads either storage format, shards the parameters over the visible TPU devices, and decodes with a KV cache:

python -m pip install --upgrade "jax[tpu]"
python -m pip install -r jax_runtime/requirements.txt
python -m jax_runtime.smoke_generate \
  --folder qwen3-1.7b-expdis-single-round \
  --max-new-tokens 16

It is meant for checking that a checkpoint loads and for small evaluations, not for high-throughput serving. See jax_runtime/README.md.

Evaluation settings

The paper evaluates with the following settings:

  • Prompt: the system message Please reason step by step, and put your final answer within \boxed{}. and the question as the user message, in the checkpoint's own chat template, with thinking enabled for Qwen3.
  • Sampling: temperature 0.6, top-p 0.95, top-k 20, min-p 0, at most 32,768 completion tokens.
  • 64 samples per problem (32 for AMC23 and 8 for GSM8K). pass@k is the unbiased estimator of Chen et al. (2021), 1 - C(n - c, k) / C(n, k) for c correct samples out of n, averaged over problems.
  • Benchmarks: AIME24, AIME25, AIME26, MATH500, and Minerva-Math (the reported mean), plus AMC23 and GSM8K.

License

Apache License 2.0. The base models, Qwen3 and Ministral 3, are also released under Apache 2.0.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support