Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
49 commits
Select commit Hold shift + click to select a range
30f94b8
Add Ministral3 EAGLE3 support
andylolu2 Jul 2, 2026
65e0931
Add Ministral3 EAGLE3 smoke scripts
andylolu2 Jul 2, 2026
adb7abb
Add Transformers data generation path
andylolu2 Jul 2, 2026
c7407c3
Add Ministral3 10k reproduction pipeline
andylolu2 Jul 2, 2026
029d5f2
Add accelerate for FP8 target loading
andylolu2 Jul 2, 2026
05da4f3
Tighten Ministral3 smoke job resources
andylolu2 Jul 2, 2026
38ba345
Shorten Ministral3 generation walltime
andylolu2 Jul 2, 2026
c09f9e2
Add kernels for FP8 target execution
andylolu2 Jul 2, 2026
87904b5
Disable hub kernels for FP8 compatibility
andylolu2 Jul 2, 2026
7e28309
Handle zero-worker training dataloaders
andylolu2 Jul 2, 2026
6486e52
Write eval metrics JSON
andylolu2 Jul 2, 2026
6d22729
Add save-only training checkpoints
andylolu2 Jul 3, 2026
248d853
Add Ministral3 training metric plots
andylolu2 Jul 3, 2026
5d26429
Tighten checkpoint cadence for preemptions
andylolu2 Jul 3, 2026
181b0e8
Add segmented Ministral3 training job
andylolu2 Jul 3, 2026
1c418bb
Add eval acceptance comparison plot
andylolu2 Jul 3, 2026
5d1a95a
Write eval metrics incrementally
andylolu2 Jul 4, 2026
e04b1f8
Add split eval task launcher
andylolu2 Jul 4, 2026
6137436
Add Ministral3 reproduction report
andylolu2 Jul 4, 2026
5ab02fb
Add reproduction report plots
andylolu2 Jul 4, 2026
ceb810f
Add Eagle3 e2e TV loss mode
andylolu2 Jul 6, 2026
7a44789
Add Ministral3 DSpark support
andylolu2 Jul 6, 2026
1a2e973
Add Ministral3 follow-up experiment launchers
andylolu2 Jul 6, 2026
613c107
Shorten follow-up training wall times
andylolu2 Jul 6, 2026
0303186
Use tighter follow-up training backfill windows
andylolu2 Jul 6, 2026
54fcc22
Add single-node follow-up fallback jobs
andylolu2 Jul 6, 2026
08823ab
Add fallback eval launchers
andylolu2 Jul 6, 2026
38a0dd6
Add TTT5 e2e TV follow-up variant
andylolu2 Jul 6, 2026
099c54e
Add follow-up training plot helper
andylolu2 Jul 6, 2026
b01ddc4
Add KL warm-start e2e TV variant
andylolu2 Jul 6, 2026
309d01a
Add lower LR KL warm-start TV variant
andylolu2 Jul 6, 2026
66bf010
Add 8GPU continuation for low LR TV warm start
andylolu2 Jul 6, 2026
ecadc74
Remove incompatible 8GPU TV continuation
andylolu2 Jul 6, 2026
e6f1dd9
Add fresh 8GPU low LR TV warm start
andylolu2 Jul 6, 2026
9b40b46
Shorten 8GPU TV warm start slices
andylolu2 Jul 6, 2026
b5287f8
Use shorter 8GPU TV warm start slices
andylolu2 Jul 6, 2026
cb37767
Shorten 8GPU TV warm-start slices
andylolu2 Jul 6, 2026
7b53e77
Shorten DSpark continuation slices
andylolu2 Jul 6, 2026
f948c49
Use shorter DSpark continuation slices
andylolu2 Jul 8, 2026
c5c0443
Use backfill-friendly DSpark slices
andylolu2 Jul 8, 2026
f6f6b78
Checkpoint DSpark more frequently
andylolu2 Jul 8, 2026
6a1dcda
Shorten DSpark eval walltime
andylolu2 Jul 8, 2026
d35c26b
Add arena-only DSpark eval rescue
andylolu2 Jul 8, 2026
bc2447b
Add DSpark eval merge helper
andylolu2 Jul 8, 2026
7a8b2ec
Shorten arena rescue eval walltime
andylolu2 Jul 8, 2026
d4a1c0c
Add four-GPU arena eval fallback
andylolu2 Jul 8, 2026
f66d102
Add two-GPU arena eval fallback
andylolu2 Jul 8, 2026
34451ed
Report Ministral DSpark follow-up results
andylolu2 Jul 9, 2026
ae8d73c
Add Ministral follow-up report plots
andylolu2 Jul 9, 2026
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
65 changes: 65 additions & 0 deletions config/dspark/dspark_ministral3_3b.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
import os

from deepspec.trainer import Ministral3DSparkTrainer

BASE_TB_DIR = "/mnt/vast/runs/andy/deepspec_tensorboard"
BASE_CKPT_DIR = "/mnt/vast/runs/andy/deepspec_checkpoints"
project_name = "deepspec"
exp_name = "dspark_block7_ministral3_3b"
seed = 42

model = dict(
target_model_name_or_path="mistralai/Ministral-3-3B-Instruct-2512",
block_size=7,
num_draft_layers=5,
target_layer_ids=[1, 7, 13, 19, 24],
mask_token_id=11,
num_anchors=512,
markov_rank=256,
markov_head_type="vanilla",
confidence_head_alpha=1.0,
confidence_head_with_markov=True,
loss_decay_gamma=4.0,
ce_loss_alpha=0.1,
l1_loss_alpha=0.9,
)

train = dict(
trainer_cls=Ministral3DSparkTrainer,
lr=6.0e-4,
warmup_ratio=0.04,
weight_decay=0.0,
precision="bf16",
local_batch_size=1,
global_batch_size=512,
num_train_epochs=10,
max_train_steps=None,
max_grad_norm=1.0,
sharding_strategy="no_shard",
torch_compile=False,
)

logging = dict(
logging_steps=10,
checkpointing_steps=3000,
save_only_checkpointing_steps=None,
keep_last_checkpoints=None,
)

data = dict(
target_cache_path=None,
chat_template="ministral3",
max_length=4096,
num_workers=4,
)


def finalize_cfg(cfg):
logging_cfg = dict(cfg["logging"])
project_name = str(cfg["project_name"])
exp_name = str(cfg["exp_name"])
logging_cfg["checkpoint_dir"] = os.path.join(BASE_CKPT_DIR, project_name, exp_name)
logging_cfg["tensorboard_dir"] = os.path.join(BASE_TB_DIR, project_name, exp_name)
cfg["logging"] = logging_cfg

return cfg
58 changes: 58 additions & 0 deletions config/eagle3/eagle3_ministral3_3b.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
import os

from deepspec.trainer import Ministral3Eagle3Trainer

BASE_TB_DIR = os.path.expanduser("~/tensorboard")
BASE_CKPT_DIR = os.path.expanduser("~/checkpoints")
project_name = "deepspec"
exp_name = "eagle3_ttt7_ministral3_3b"
seed = 0

model = dict(
target_model_name_or_path="mistralai/Ministral-3-3B-Instruct-2512",
target_layer_ids=[1, 7, 13, 19, 24],
ttt_length=7,
step_loss_decay=0.8,
loss_type="soft_ce",
draft_num_hidden_layers=1,
)

train = dict(
trainer_cls=Ministral3Eagle3Trainer,
lr=6.0e-4,
warmup_ratio=0.04,
weight_decay=0.0,
precision="bf16",
local_batch_size=1,
global_batch_size=512,
num_train_epochs=10,
max_train_steps=None,
max_grad_norm=1.0,
sharding_strategy="no_shard",
torch_compile=False,
)

logging = dict(
logging_steps=10,
checkpointing_steps=3000,
save_only_checkpointing_steps=None,
keep_last_checkpoints=None,
)

data = dict(
target_cache_path=None,
chat_template="ministral3",
max_length=4096,
num_workers=4,
)


def finalize_cfg(cfg):
logging_cfg = dict(cfg["logging"])
project_name = str(cfg["project_name"])
exp_name = str(cfg["exp_name"])
logging_cfg["checkpoint_dir"] = os.path.join(BASE_CKPT_DIR, project_name, exp_name)
logging_cfg["tensorboard_dir"] = os.path.join(BASE_TB_DIR, project_name, exp_name)
cfg["logging"] = logging_cfg

return cfg
20 changes: 20 additions & 0 deletions config/eagle3/eagle3_ministral3_3b_e2e_tv.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
import os

from config.eagle3.eagle3_ministral3_3b import * # noqa: F403

BASE_TB_DIR = "/mnt/vast/runs/andy/deepspec_tensorboard"
BASE_CKPT_DIR = "/mnt/vast/runs/andy/deepspec_checkpoints"
exp_name = "eagle3_ttt7_ministral3_3b_e2e_tv"
model = dict(model) # noqa: F405
model["loss_type"] = "e2e_tv"


def finalize_cfg(cfg):
logging_cfg = dict(cfg["logging"])
project_name = str(cfg["project_name"])
exp_name = str(cfg["exp_name"])
logging_cfg["checkpoint_dir"] = os.path.join(BASE_CKPT_DIR, project_name, exp_name)
logging_cfg["tensorboard_dir"] = os.path.join(BASE_TB_DIR, project_name, exp_name)
cfg["logging"] = logging_cfg

return cfg
26 changes: 26 additions & 0 deletions config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import os

from config.eagle3.eagle3_ministral3_3b import * # noqa: F403

BASE_TB_DIR = "/mnt/vast/runs/andy/deepspec_tensorboard"
BASE_CKPT_DIR = "/mnt/vast/runs/andy/deepspec_checkpoints"
KL_DRAFT_CHECKPOINT = (
"/mnt/vast/home/andy/checkpoints/deepspec/"
"eagle3_ttt7_ministral3_3b_10k_3000steps/step_latest"
)

exp_name = "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl"
model = dict(model) # noqa: F405
model["loss_type"] = "e2e_tv"
model["init_draft_model_name_or_path"] = KL_DRAFT_CHECKPOINT


def finalize_cfg(cfg):
logging_cfg = dict(cfg["logging"])
project_name = str(cfg["project_name"])
exp_name = str(cfg["exp_name"])
logging_cfg["checkpoint_dir"] = os.path.join(BASE_CKPT_DIR, project_name, exp_name)
logging_cfg["tensorboard_dir"] = os.path.join(BASE_TB_DIR, project_name, exp_name)
cfg["logging"] = logging_cfg

return cfg
5 changes: 5 additions & 0 deletions config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
from config.eagle3.eagle3_ministral3_3b_e2e_tv_from_kl import * # noqa: F403

exp_name = "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4"
train = dict(train) # noqa: F405
train["lr"] = 1.0e-4
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from config.eagle3.eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4 import * # noqa: F403

exp_name = "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu"
21 changes: 21 additions & 0 deletions config/eagle3/eagle3_ministral3_3b_e2e_tv_ttt5.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
import os

from config.eagle3.eagle3_ministral3_3b import * # noqa: F403

BASE_TB_DIR = "/mnt/vast/runs/andy/deepspec_tensorboard"
BASE_CKPT_DIR = "/mnt/vast/runs/andy/deepspec_checkpoints"
exp_name = "eagle3_ttt5_ministral3_3b_e2e_tv"
model = dict(model) # noqa: F405
model["loss_type"] = "e2e_tv"
model["ttt_length"] = 5


def finalize_cfg(cfg):
logging_cfg = dict(cfg["logging"])
project_name = str(cfg["project_name"])
exp_name = str(cfg["exp_name"])
logging_cfg["checkpoint_dir"] = os.path.join(BASE_CKPT_DIR, project_name, exp_name)
logging_cfg["tensorboard_dir"] = os.path.join(BASE_TB_DIR, project_name, exp_name)
cfg["logging"] = logging_cfg

return cfg
10 changes: 10 additions & 0 deletions deepspec/data/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,16 @@ def get(self, name):
),
)

TEMPLATE_REGISTRY.register(
"ministral3",
ChatTemplate(
assistant_header="[/INST]",
user_header="[INST]",
system_prompt=None,
end_of_turn_token="</s>",
),
)


class GeneralParser:
def __init__(self, tokenizer, chat_template):
Expand Down
14 changes: 12 additions & 2 deletions deepspec/eval/__init__.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,22 @@
from .base_evaluator import BaseEvaluator, DraftProposal, VerificationResult
from .dspark import Gemma4DSparkEvaluator, Qwen3DSparkEvaluator
from .eagle3 import Gemma4Eagle3Evaluator, Qwen3Eagle3Evaluator
from .dspark import (
Gemma4DSparkEvaluator,
Ministral3DSparkEvaluator,
Qwen3DSparkEvaluator,
)
from .eagle3 import (
Gemma4Eagle3Evaluator,
Ministral3Eagle3Evaluator,
Qwen3Eagle3Evaluator,
)

__all__ = [
"BaseEvaluator",
"DraftProposal",
"Gemma4Eagle3Evaluator",
"Gemma4DSparkEvaluator",
"Ministral3DSparkEvaluator",
"Ministral3Eagle3Evaluator",
"Qwen3Eagle3Evaluator",
"Qwen3DSparkEvaluator",
"VerificationResult",
Expand Down
26 changes: 26 additions & 0 deletions deepspec/eval/base_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -702,12 +702,38 @@ def record_dataset_metrics(
)
self.metrics_rows.append(metrics_row)
self.print_dataset_result(metrics_row)
self.write_output_json(complete=False)
return metrics_row

def write_output_json(self, *, complete: bool) -> None:
if (
self.args.output_json is None
or dist.get_rank() != 0
or not self.metrics_rows
):
return
output_path = Path(self.args.output_json)
output_path.parent.mkdir(parents=True, exist_ok=True)
tmp_path = output_path.with_suffix(output_path.suffix + ".tmp")
with tmp_path.open("w", encoding="utf-8") as handle:
json.dump(
{
"target_model": self.args.target_name_or_path,
"draft_model": self.args.draft_name_or_path,
"step": self.args.step,
"complete": complete,
"rows": self.metrics_rows,
},
handle,
indent=2,
)
tmp_path.replace(output_path)

def report_results(self) -> None:
if dist.get_rank() == 0 and self.metrics_rows:
if self.args.tensorboard_dir is not None:
self.log_tensorboard()
self.write_output_json(complete=True)
self.print_results()

def evaluate(self) -> None:
Expand Down
7 changes: 6 additions & 1 deletion deepspec/eval/dspark/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,8 @@
from .evaluator import Gemma4DSparkEvaluator, Qwen3DSparkEvaluator
from .evaluator import (
Gemma4DSparkEvaluator,
Ministral3DSparkEvaluator,
Qwen3DSparkEvaluator,
)
from .draft_ops import (
DSparkDraftProposal,
build_dspark_proposal,
Expand All @@ -8,6 +12,7 @@

__all__ = [
"Gemma4DSparkEvaluator",
"Ministral3DSparkEvaluator",
"Qwen3DSparkEvaluator",
"DSparkDraftProposal",
"build_dspark_proposal",
Expand Down
3 changes: 2 additions & 1 deletion deepspec/eval/dspark/draft_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,11 @@
from deepspec.eval.base_evaluator import DraftProposal
from deepspec.utils.sampling import logits_to_probs
from deepspec.modeling.dspark.gemma4 import Gemma4DSparkModel
from deepspec.modeling.dspark.ministral3 import Ministral3DSparkModel
from deepspec.modeling.dspark.qwen3 import Qwen3DSparkModel


DSparkModel = Qwen3DSparkModel | Gemma4DSparkModel
DSparkModel = Qwen3DSparkModel | Gemma4DSparkModel | Ministral3DSparkModel


@dataclass
Expand Down
12 changes: 9 additions & 3 deletions deepspec/eval/dspark/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from types import SimpleNamespace

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, DynamicCache
from transformers import AutoTokenizer, DynamicCache

from deepspec.eval.base_evaluator import (
BaseEvaluator,
Expand All @@ -21,7 +21,9 @@
)
from deepspec.modeling.dspark.common import extract_context_feature
from deepspec.modeling.dspark.gemma4 import Gemma4DSparkModel
from deepspec.modeling.dspark.ministral3 import Ministral3DSparkModel
from deepspec.modeling.dspark.qwen3 import Qwen3DSparkModel
from deepspec.modeling.target_utils import load_target_causal_lm, load_target_tokenizer
from deepspec.utils import jsonable


Expand Down Expand Up @@ -66,7 +68,7 @@ def _build_confidence_head_recorder(self) -> ConfidenceHeadRecorder | None:
)

def build_models(self) -> tuple[object, Qwen3DSparkModel, AutoTokenizer]:
target_model = AutoModelForCausalLM.from_pretrained(
target_model = load_target_causal_lm(
self.args.target_name_or_path,
dtype=torch.bfloat16,
attn_implementation=self.EVAL_ATTN_IMPLEMENTATION,
Expand All @@ -79,7 +81,7 @@ def build_models(self) -> tuple[object, Qwen3DSparkModel, AutoTokenizer]:
).to(self.device).eval()
assert_no_final_target_layer(target_model, draft_model.target_layer_ids)
assert 0.0 <= float(self.args.confidence_threshold) <= 1.0
tokenizer = AutoTokenizer.from_pretrained(self.args.target_name_or_path)
tokenizer = load_target_tokenizer(self.args.target_name_or_path)
return target_model, draft_model, tokenizer

def _init_context(
Expand Down Expand Up @@ -223,3 +225,7 @@ def print_results(self) -> None:

class Gemma4DSparkEvaluator(Qwen3DSparkEvaluator):
draft_model_cls = Gemma4DSparkModel


class Ministral3DSparkEvaluator(Qwen3DSparkEvaluator):
draft_model_cls = Ministral3DSparkModel
12 changes: 10 additions & 2 deletions deepspec/eval/eagle3/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,11 @@
from .evaluator import Gemma4Eagle3Evaluator, Qwen3Eagle3Evaluator
from .evaluator import (
Gemma4Eagle3Evaluator,
Ministral3Eagle3Evaluator,
Qwen3Eagle3Evaluator,
)

__all__ = ["Gemma4Eagle3Evaluator", "Qwen3Eagle3Evaluator"]
__all__ = [
"Gemma4Eagle3Evaluator",
"Ministral3Eagle3Evaluator",
"Qwen3Eagle3Evaluator",
]
Loading