diff --git a/config/dspark/dspark_ministral3_3b.py b/config/dspark/dspark_ministral3_3b.py new file mode 100644 index 00000000..35ad9c26 --- /dev/null +++ b/config/dspark/dspark_ministral3_3b.py @@ -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 diff --git a/config/eagle3/eagle3_ministral3_3b.py b/config/eagle3/eagle3_ministral3_3b.py new file mode 100644 index 00000000..6f2a8de4 --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b.py @@ -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 diff --git a/config/eagle3/eagle3_ministral3_3b_e2e_tv.py b/config/eagle3/eagle3_ministral3_3b_e2e_tv.py new file mode 100644 index 00000000..cf15a444 --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b_e2e_tv.py @@ -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 diff --git a/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl.py b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl.py new file mode 100644 index 00000000..95431bc0 --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl.py @@ -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 diff --git a/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4.py b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4.py new file mode 100644 index 00000000..3ebdbfe6 --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4.py @@ -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 diff --git a/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu.py b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu.py new file mode 100644 index 00000000..fb19c32b --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu.py @@ -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" diff --git a/config/eagle3/eagle3_ministral3_3b_e2e_tv_ttt5.py b/config/eagle3/eagle3_ministral3_3b_e2e_tv_ttt5.py new file mode 100644 index 00000000..e58720fd --- /dev/null +++ b/config/eagle3/eagle3_ministral3_3b_e2e_tv_ttt5.py @@ -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 diff --git a/deepspec/data/parser.py b/deepspec/data/parser.py index e7828dc8..127ec167 100644 --- a/deepspec/data/parser.py +++ b/deepspec/data/parser.py @@ -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="", + ), +) + class GeneralParser: def __init__(self, tokenizer, chat_template): diff --git a/deepspec/eval/__init__.py b/deepspec/eval/__init__.py index 8eaead30..8f74026a 100644 --- a/deepspec/eval/__init__.py +++ b/deepspec/eval/__init__.py @@ -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", diff --git a/deepspec/eval/base_evaluator.py b/deepspec/eval/base_evaluator.py index d094ec38..e9451b96 100644 --- a/deepspec/eval/base_evaluator.py +++ b/deepspec/eval/base_evaluator.py @@ -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: diff --git a/deepspec/eval/dspark/__init__.py b/deepspec/eval/dspark/__init__.py index 2126a6cc..6bfd295e 100644 --- a/deepspec/eval/dspark/__init__.py +++ b/deepspec/eval/dspark/__init__.py @@ -1,4 +1,8 @@ -from .evaluator import Gemma4DSparkEvaluator, Qwen3DSparkEvaluator +from .evaluator import ( + Gemma4DSparkEvaluator, + Ministral3DSparkEvaluator, + Qwen3DSparkEvaluator, +) from .draft_ops import ( DSparkDraftProposal, build_dspark_proposal, @@ -8,6 +12,7 @@ __all__ = [ "Gemma4DSparkEvaluator", + "Ministral3DSparkEvaluator", "Qwen3DSparkEvaluator", "DSparkDraftProposal", "build_dspark_proposal", diff --git a/deepspec/eval/dspark/draft_ops.py b/deepspec/eval/dspark/draft_ops.py index 753a56d9..42f3cb01 100644 --- a/deepspec/eval/dspark/draft_ops.py +++ b/deepspec/eval/dspark/draft_ops.py @@ -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 diff --git a/deepspec/eval/dspark/evaluator.py b/deepspec/eval/dspark/evaluator.py index eba2b34d..3ffa3ac3 100644 --- a/deepspec/eval/dspark/evaluator.py +++ b/deepspec/eval/dspark/evaluator.py @@ -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, @@ -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 @@ -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, @@ -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( @@ -223,3 +225,7 @@ def print_results(self) -> None: class Gemma4DSparkEvaluator(Qwen3DSparkEvaluator): draft_model_cls = Gemma4DSparkModel + + +class Ministral3DSparkEvaluator(Qwen3DSparkEvaluator): + draft_model_cls = Ministral3DSparkModel diff --git a/deepspec/eval/eagle3/__init__.py b/deepspec/eval/eagle3/__init__.py index 7765f306..2bd14ada 100644 --- a/deepspec/eval/eagle3/__init__.py +++ b/deepspec/eval/eagle3/__init__.py @@ -1,3 +1,11 @@ -from .evaluator import Gemma4Eagle3Evaluator, Qwen3Eagle3Evaluator +from .evaluator import ( + Gemma4Eagle3Evaluator, + Ministral3Eagle3Evaluator, + Qwen3Eagle3Evaluator, +) -__all__ = ["Gemma4Eagle3Evaluator", "Qwen3Eagle3Evaluator"] +__all__ = [ + "Gemma4Eagle3Evaluator", + "Ministral3Eagle3Evaluator", + "Qwen3Eagle3Evaluator", +] diff --git a/deepspec/eval/eagle3/evaluator.py b/deepspec/eval/eagle3/evaluator.py index 1ddb9cae..b2c68b1a 100644 --- a/deepspec/eval/eagle3/evaluator.py +++ b/deepspec/eval/eagle3/evaluator.py @@ -3,7 +3,7 @@ from types import SimpleNamespace import torch -from transformers import AutoModelForCausalLM, AutoTokenizer, DynamicCache +from transformers import DynamicCache from deepspec.eval.base_evaluator import ( BaseEvaluator, @@ -15,7 +15,9 @@ ) from deepspec.modeling.eagle3 import extract_eagle3_context_feature from deepspec.modeling.eagle3.gemma4 import Gemma4Eagle3Model +from deepspec.modeling.eagle3.ministral3 import Ministral3Eagle3Model from deepspec.modeling.eagle3.qwen3 import Qwen3Eagle3Model +from deepspec.modeling.target_utils import load_target_causal_lm, load_target_tokenizer from deepspec.utils.sampling import logits_to_probs, sample_tokens @@ -40,8 +42,8 @@ def __init__(self, local_rank: int, args): def max_proposal_tokens(self) -> int: return int(self.draft_model.ttt_length) - def build_models(self) -> tuple[object, Qwen3Eagle3Model, AutoTokenizer]: - target_model = AutoModelForCausalLM.from_pretrained( + def build_models(self) -> tuple[object, Qwen3Eagle3Model, object]: + target_model = load_target_causal_lm( self.args.target_name_or_path, dtype=torch.bfloat16, attn_implementation=self.EVAL_ATTN_IMPLEMENTATION, @@ -55,7 +57,7 @@ def build_models(self) -> tuple[object, Qwen3Eagle3Model, AutoTokenizer]: draft_model.target_layer_ids = [int(x) for x in draft_model.target_layer_ids] assert_no_final_target_layer(target_model, draft_model.target_layer_ids) - 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( @@ -192,4 +194,12 @@ class Gemma4Eagle3Evaluator(Qwen3Eagle3Evaluator): draft_model_cls = Gemma4Eagle3Model -__all__ = ["Gemma4Eagle3Evaluator", "Qwen3Eagle3Evaluator"] +class Ministral3Eagle3Evaluator(Qwen3Eagle3Evaluator): + draft_model_cls = Ministral3Eagle3Model + + +__all__ = [ + "Gemma4Eagle3Evaluator", + "Ministral3Eagle3Evaluator", + "Qwen3Eagle3Evaluator", +] diff --git a/deepspec/modeling/__init__.py b/deepspec/modeling/__init__.py index 22df914e..81637fcd 100644 --- a/deepspec/modeling/__init__.py +++ b/deepspec/modeling/__init__.py @@ -1,6 +1,11 @@ +import os + +os.environ["USE_HUB_KERNELS"] = "0" + from .dspark import ( DSparkForwardOutput, Gemma4DSparkModel, + Ministral3DSparkModel, Qwen3DSparkModel, ) from .eagle3 import Gemma4Eagle3Model, Qwen3Eagle3Model @@ -9,6 +14,7 @@ "DSparkForwardOutput", "Gemma4Eagle3Model", "Gemma4DSparkModel", + "Ministral3DSparkModel", "Qwen3Eagle3Model", "Qwen3DSparkModel", ] diff --git a/deepspec/modeling/dspark/__init__.py b/deepspec/modeling/dspark/__init__.py index 723b008d..ee2473fd 100644 --- a/deepspec/modeling/dspark/__init__.py +++ b/deepspec/modeling/dspark/__init__.py @@ -1,10 +1,12 @@ from .common import DSparkForwardOutput, extract_context_feature from .gemma4 import Gemma4DSparkModel +from .ministral3 import Ministral3DSparkModel from .qwen3 import Qwen3DSparkModel __all__ = [ "DSparkForwardOutput", "extract_context_feature", "Gemma4DSparkModel", + "Ministral3DSparkModel", "Qwen3DSparkModel", ] diff --git a/deepspec/modeling/dspark/ministral3/__init__.py b/deepspec/modeling/dspark/ministral3/__init__.py new file mode 100644 index 00000000..70ed1c50 --- /dev/null +++ b/deepspec/modeling/dspark/ministral3/__init__.py @@ -0,0 +1,5 @@ +from .modeling import Ministral3DSparkModel + +__all__ = [ + "Ministral3DSparkModel", +] diff --git a/deepspec/modeling/dspark/ministral3/config.py b/deepspec/modeling/dspark/ministral3/config.py new file mode 100644 index 00000000..8d49a19e --- /dev/null +++ b/deepspec/modeling/dspark/ministral3/config.py @@ -0,0 +1,73 @@ +import copy + +from deepspec.modeling.dspark.common import validate_target_layer_ids + + +TRAIN_ATTN_IMPLEMENTATION = "flex_attention" + + +def build_draft_config( + target_config, + model_args, +): + assert target_config.model_type == "mistral3", ( + "Ministral3 DSpark expects the top-level Mistral3 target config, " + f"got model_type={target_config.model_type!r}." + ) + text_config = copy.deepcopy(target_config.text_config) + assert text_config.model_type == "ministral3", ( + "Ministral3 DSpark expects target_config.text_config.model_type to be " + f"'ministral3', got {text_config.model_type!r}." + ) + num_target_layers = int(text_config.num_hidden_layers) + num_draft_layers = int(model_args.num_draft_layers) + layer_types = ["full_attention"] * num_draft_layers + assert "target_layer_ids" in model_args, "target_layer_ids must be provided." + target_layer_ids = validate_target_layer_ids( + model_args.target_layer_ids, + num_target_layers, + ) + + confidence_head_alpha = float(model_args.confidence_head_alpha) + assert confidence_head_alpha >= 0.0 + enable_confidence_head = confidence_head_alpha > 0.0 + if enable_confidence_head: + assert "confidence_head_with_markov" in model_args, ( + "confidence_head_with_markov must be provided when " + "confidence_head_alpha > 0." + ) + markov_rank = int(model_args.markov_rank) + assert markov_rank >= 0, f"markov_rank must be >= 0, got {markov_rank}" + if markov_rank > 0: + assert "markov_head_type" in model_args, ( + "markov_head_type must be provided when markov_rank > 0." + ) + + draft_config = text_config + draft_config.architectures = ["Ministral3DSparkModel"] + draft_config.target_model_type = str(target_config.model_type) + draft_config.target_text_model_type = str(text_config.model_type) + draft_config.target_model_name_or_path = str(model_args.target_model_name_or_path) + draft_config.num_target_layers = num_target_layers + draft_config.num_hidden_layers = num_draft_layers + draft_config.block_size = int(model_args.block_size) + draft_config.tie_word_embeddings = False + draft_config.layer_types = layer_types + draft_config._attn_implementation = TRAIN_ATTN_IMPLEMENTATION + draft_config.mask_token_id = int(model_args.mask_token_id) + draft_config.target_layer_ids = target_layer_ids + draft_config.num_anchors = int(model_args.num_anchors) + draft_config.enable_confidence_head = enable_confidence_head + if enable_confidence_head: + draft_config.confidence_head_with_markov = bool( + model_args.confidence_head_with_markov + ) + draft_config.markov_rank = markov_rank + if markov_rank > 0: + draft_config.markov_head_type = str(model_args.markov_head_type) + return draft_config + + +__all__ = [ + "build_draft_config", +] diff --git a/deepspec/modeling/dspark/ministral3/modeling.py b/deepspec/modeling/dspark/ministral3/modeling.py new file mode 100644 index 00000000..9c5ac108 --- /dev/null +++ b/deepspec/modeling/dspark/ministral3/modeling.py @@ -0,0 +1,547 @@ +from __future__ import annotations + +from typing import Callable, Optional + +import torch +from torch import nn + +from transformers.cache_utils import Cache +from transformers.models.ministral3.modeling_ministral3 import ( + ALL_ATTENTION_FUNCTIONS, + FlashAttentionKwargs, + GradientCheckpointingLayer, + Ministral3MLP, + Ministral3PreTrainedModel, + Ministral3RMSNorm, + Ministral3RotaryEmbedding, + eager_attention_forward, + get_llama_4_attn_scale, + rotate_half, +) +from typing_extensions import Tuple, Unpack + +from deepspec.modeling.dspark.common import ( + AcceptRatePredictor, + DSparkForwardOutput, + build_eval_mask, + create_dspark_attention_mask, + create_noise_embed, + create_position_ids, + log_sampler_stats, + sample_anchor_positions, +) +from deepspec.modeling.dspark.markov_head import build_markov_head +from deepspec.utils.sampling import sample_tokens + + +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + q_len = q.size(-2) + q_embed = (q * cos[..., -q_len:, :]) + (rotate_half(q) * sin[..., -q_len:, :]) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed + + +class Ministral3DSparkAttention(nn.Module): + def __init__(self, config, layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.head_dim = getattr( + config, "head_dim", config.hidden_size // config.num_attention_heads + ) + self.num_attention_heads = config.num_attention_heads + self.num_key_value_heads = config.num_key_value_heads + self.num_key_value_groups = ( + self.num_attention_heads // self.num_key_value_heads + ) + self.scaling = self.head_dim**-0.5 + self.attention_dropout = config.attention_dropout + self.is_causal = False + attention_bias = bool(getattr(config, "attention_bias", False)) + self.q_proj = nn.Linear( + config.hidden_size, + self.num_attention_heads * self.head_dim, + bias=attention_bias, + ) + self.k_proj = nn.Linear( + config.hidden_size, + self.num_key_value_heads * self.head_dim, + bias=attention_bias, + ) + self.v_proj = nn.Linear( + config.hidden_size, + self.num_key_value_heads * self.head_dim, + bias=attention_bias, + ) + self.o_proj = nn.Linear( + self.num_attention_heads * self.head_dim, + config.hidden_size, + bias=attention_bias, + ) + self.sliding_window = ( + config.sliding_window + if config.layer_types[layer_idx] == "sliding_attention" + else None + ) + + def forward( + self, + hidden_states: torch.Tensor, + target_hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: Optional[torch.Tensor], + position_ids: torch.LongTensor, + past_key_values: Optional[Cache] = None, + cache_position: Optional[torch.LongTensor] = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + bsz, q_len = hidden_states.shape[:-1] + ctx_len = target_hidden_states.shape[1] + q = self.q_proj(hidden_states).view( + bsz, q_len, self.num_attention_heads, self.head_dim + ) + q = q.transpose(1, 2) + k_ctx = self.k_proj(target_hidden_states) + k_noise = self.k_proj(hidden_states) + v_ctx = self.v_proj(target_hidden_states) + v_noise = self.v_proj(hidden_states) + k = torch.cat([k_ctx, k_noise], dim=1).view( + bsz, ctx_len + q_len, self.num_key_value_heads, self.head_dim + ) + v = torch.cat([v_ctx, v_noise], dim=1).view( + bsz, ctx_len + q_len, self.num_key_value_heads, self.head_dim + ) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + q_position_ids = position_ids[:, -q_len:] + q = q * get_llama_4_attn_scale( + q_position_ids, + self.config.rope_parameters.get("llama_4_scaling_beta"), + self.config.rope_parameters.get("original_max_position_embeddings"), + ).to(q.dtype) + if past_key_values is not None: + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + k, v = past_key_values.update(k, v, self.layer_idx, cache_kwargs) + if ( + self.config._attn_implementation == "flex_attention" + and self.num_key_value_groups > 1 + ): + kv_seq_len = k.shape[-2] + k = k.repeat_interleave(self.num_key_value_groups, dim=1) + v = v.repeat_interleave(self.num_key_value_groups, dim=1) + k = k.reshape(bsz, self.num_attention_heads, kv_seq_len, self.head_dim) + v = v.reshape(bsz, self.num_attention_heads, kv_seq_len, self.head_dim) + attn_fn: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, + eager_attention_forward, + ) + attn_is_causal = bool(kwargs.get("is_causal", False)) + # The SDPA path may consult module.is_causal when dispatching kernels, + # so keep the per-call value mirrored on the module before invoking it. + self.is_causal = attn_is_causal + kwargs["is_causal"] = attn_is_causal + attn_output, attn_weights = attn_fn( + self, + q, + k, + v, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + sliding_window=self.sliding_window, + **kwargs, + ) + attn_output = attn_output.reshape( + bsz, q_len, self.num_attention_heads * self.head_dim + ) + return self.o_proj(attn_output), attn_weights + + +class Ministral3DSparkDecoderLayer(GradientCheckpointingLayer): + def __init__(self, config, layer_idx: int): + super().__init__() + self.hidden_size = config.hidden_size + self.self_attn = Ministral3DSparkAttention(config=config, layer_idx=layer_idx) + self.mlp = Ministral3MLP(config) + self.input_layernorm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = Ministral3RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + def forward( + self, + target_hidden_states: Optional[torch.Tensor] = None, + hidden_states: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_value: Optional[Cache] = None, + output_attentions: Optional[bool] = False, + use_cache: Optional[bool] = False, + cache_position: Optional[torch.LongTensor] = None, + position_embeddings: Optional[ + Tuple[torch.Tensor, torch.Tensor] + ] = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + hidden_states = self.self_attn( + hidden_states=hidden_states, + target_hidden_states=target_hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_value, + output_attentions=output_attentions, + use_cache=use_cache, + cache_position=cache_position, + position_embeddings=position_embeddings, + **kwargs, + )[0] + hidden_states = residual + hidden_states + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + return residual + hidden_states + + +class Ministral3DSparkModel(Ministral3PreTrainedModel): + _no_split_modules = ["Ministral3DSparkDecoderLayer"] + + def __init__(self, config) -> None: + super().__init__(config) + self.config = config + required_fields = ( + "target_layer_ids", + "mask_token_id", + "num_anchors", + "enable_confidence_head", + "markov_rank", + ) + for field in required_fields: + assert hasattr(config, field), f"config.{field} must be provided." + if int(config.markov_rank) > 0: + assert hasattr(config, "markov_head_type"), ( + "config.markov_head_type must be provided when markov_rank > 0." + ) + if bool(config.enable_confidence_head): + assert hasattr(config, "confidence_head_with_markov"), ( + "config.confidence_head_with_markov must be provided when " + "enable_confidence_head is true." + ) + self.target_layer_ids = config.target_layer_ids + + self.embed_tokens = nn.Embedding( + config.vocab_size, + config.hidden_size, + padding_idx=getattr(config, "pad_token_id", None), + ) + self.layers = nn.ModuleList( + [ + Ministral3DSparkDecoderLayer(config, layer_idx) + for layer_idx in range(config.num_hidden_layers) + ] + ) + self.norm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.rotary_emb = Ministral3RotaryEmbedding(config) + self.fc = nn.Linear( + len(self.target_layer_ids) * config.hidden_size, + config.hidden_size, + bias=False, + ) + self.hidden_norm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) + self.block_size = int(config.block_size) + self.mask_token_id = config.mask_token_id + self.num_anchors = int(config.num_anchors) + + # Markov head. + self.markov_head = build_markov_head(config) + + # Confidence head. + self.enable_confidence_head = bool(config.enable_confidence_head) + self.confidence_head_with_markov = False + if self.enable_confidence_head: + self.confidence_head_with_markov = bool(config.confidence_head_with_markov) + if self.enable_confidence_head and self.confidence_head_with_markov: + assert self.markov_head is not None + + self.confidence_head = None + if self.enable_confidence_head: + input_dim = int(config.hidden_size) + if self.confidence_head_with_markov: + input_dim += config.markov_rank + self.confidence_head = AcceptRatePredictor(input_dim=input_dim) + self.post_init() + + def initialize_embeddings_and_head( + self, + *, + embed_tokens: nn.Module, + lm_head: nn.Module, + freeze: bool = True, + ): + assert self.embed_tokens.weight.shape == embed_tokens.weight.shape + assert self.lm_head.weight.shape == lm_head.weight.shape + with torch.no_grad(): + self.embed_tokens.weight.copy_(embed_tokens.weight.detach()) + self.lm_head.weight.copy_(lm_head.weight.detach()) + if freeze: + self.set_embedding_head_trainable(False) + + def set_embedding_head_trainable(self, trainable: bool): + self.embed_tokens.requires_grad_(trainable) + self.lm_head.requires_grad_(trainable) + + def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.lm_head(hidden_states) + + def predict_confidence_step( + self, + hidden_states: torch.Tensor, + prev_token_ids: Optional[torch.Tensor] = None, + ) -> Optional[torch.Tensor]: + if self.confidence_head is None: + return None + if self.confidence_head_with_markov: + assert self.markov_head is not None + assert prev_token_ids is not None + prev_embeddings = self.markov_head.get_prev_embeddings(prev_token_ids).to( + dtype=hidden_states.dtype + ) + features = torch.cat([hidden_states, prev_embeddings], dim=-1) + return self.confidence_head(features).float() + return self.confidence_head(hidden_states).float() + + def sample_draft_tokens( + self, + base_logits: torch.Tensor, + *, + first_prev_token_ids: torch.Tensor, + temperature: float = 0.0, + hidden_states: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + batch_size, proposal_len = base_logits.shape[:2] + if proposal_len == 0: + empty_tokens = torch.empty( + batch_size, + 0, + dtype=torch.long, + device=base_logits.device, + ) + return empty_tokens, base_logits + if self.markov_head is None: + return sample_tokens(base_logits, temperature), base_logits + return self.markov_head.sample_block_tokens( + base_logits, + first_prev_token_ids=first_prev_token_ids, + hidden_states=hidden_states, + temperature=temperature, + ) + + def sample_draft_token_step( + self, + base_logits: torch.Tensor, + *, + prev_token_ids: torch.Tensor, + temperature: float = 0.0, + hidden_states: Optional[torch.Tensor] = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + assert base_logits.ndim == 2, ( + "sample_draft_token_step expects base_logits shaped [batch, vocab], " + f"got {tuple(base_logits.shape)}." + ) + if self.markov_head is None: + step_logits = base_logits + else: + step_logits = self.markov_head.apply_step_logits( + base_logits, + token_ids=prev_token_ids, + hidden_states=hidden_states, + ) + sampled_token_ids = sample_tokens( + step_logits.unsqueeze(1), + temperature=temperature, + ).squeeze(1) + return sampled_token_ids, step_logits + + def _forward_backbone( + self, + *, + position_ids: torch.LongTensor, + attention_mask: Optional[torch.Tensor] = None, + noise_embedding: Optional[torch.Tensor] = None, + target_hidden_states: Optional[torch.Tensor] = None, + past_key_values: Optional[Cache] = None, + use_cache: bool = False, + **kwargs, + ) -> torch.Tensor: + hidden_states = noise_embedding + target_hidden_states = self.hidden_norm(self.fc(target_hidden_states)) + position_embeddings = self.rotary_emb(hidden_states, position_ids) + for layer in self.layers: + hidden_states = layer( + hidden_states=hidden_states, + target_hidden_states=target_hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_value=past_key_values, + use_cache=use_cache, + position_embeddings=position_embeddings, + **kwargs, + ) + return self.norm(hidden_states) + + def forward( + self, + input_ids: torch.Tensor, + target_hidden_states: torch.Tensor, + loss_mask: torch.Tensor, + target_last_hidden_states: Optional[torch.Tensor] = None, + ) -> DSparkForwardOutput: + bsz, seq_len = input_ids.shape + device = input_ids.device + + anchor_positions, block_keep_mask = sample_anchor_positions( + seq_len=seq_len, + loss_mask=loss_mask, + num_anchors=self.num_anchors, + device=device, + ) + noise_embedding = create_noise_embed( + self.embed_tokens, + input_ids, + anchor_positions, + block_keep_mask, + mask_token_id=self.mask_token_id, + block_size=self.block_size, + ) + context_position_ids = torch.arange(seq_len, device=device).unsqueeze(0).expand( + bsz, -1 + ) + draft_position_ids = create_position_ids(anchor_positions, self.block_size) + full_position_ids = torch.cat( + [context_position_ids, draft_position_ids], + dim=1, + ) + dspark_attn_mask = create_dspark_attention_mask( + anchor_positions=anchor_positions, + block_keep_mask=block_keep_mask, + seq_len=seq_len, + block_size=self.block_size, + device=device, + ) + output_hidden = self._forward_backbone( + position_ids=full_position_ids, + noise_embedding=noise_embedding, + target_hidden_states=target_hidden_states, + attention_mask=dspark_attn_mask, + ) + + num_blocks = anchor_positions.size(1) + output_hidden_4d = output_hidden.reshape(bsz, num_blocks, self.block_size, -1) + + label_offsets = torch.arange(1, self.block_size + 1, device=device).view( + 1, 1, -1 + ) + label_indices = anchor_positions.unsqueeze(-1) + label_offsets + safe_label_indices = label_indices.clamp(max=seq_len - 1) + safe_label_indices = torch.where( + block_keep_mask.unsqueeze(-1), + safe_label_indices, + torch.zeros_like(safe_label_indices), + ) + target_ids = torch.gather( + input_ids.unsqueeze(1).expand(-1, anchor_positions.size(1), -1), + 2, + safe_label_indices, + ) + aligned_target_logits = None + if target_last_hidden_states is not None: + target_pred_indices = (safe_label_indices - 1).clamp(min=0) + aligned_target_hidden = torch.gather( + target_last_hidden_states.unsqueeze(1).expand( + -1, + anchor_positions.size(1), + -1, + -1, + ), + 2, + target_pred_indices.unsqueeze(-1).expand( + -1, + -1, + -1, + target_last_hidden_states.size(-1), + ), + ) + aligned_target_logits = self.compute_logits(aligned_target_hidden) + eval_mask = build_eval_mask( + seq_len=seq_len, + loss_mask=loss_mask, + label_indices=label_indices, + safe_label_indices=safe_label_indices, + block_keep_mask=block_keep_mask, + ) + anchor_token_ids = torch.gather( + input_ids, + 1, + anchor_positions, + ) + prev_token_ids = torch.cat( + [anchor_token_ids.unsqueeze(-1), target_ids[:, :, :-1]], + dim=-1, + ) + draft_logits = self.compute_logits(output_hidden).reshape( + bsz, + num_blocks, + self.block_size, + -1, + ) + if self.markov_head is not None: + draft_logits = self.markov_head.apply_block_logits( + draft_logits, + token_ids=prev_token_ids, + hidden_states=output_hidden_4d, + ) + + log_sampler_stats( + seq_len=seq_len, + loss_mask=loss_mask, + eval_mask=eval_mask, + block_keep_mask=block_keep_mask, + block_size=self.block_size, + num_anchors=self.num_anchors, + ) + + confidence_pred = None + if self.confidence_head is not None: + if self.confidence_head_with_markov: + prev_embeddings = self.markov_head.get_prev_embeddings(prev_token_ids).to( + dtype=output_hidden_4d.dtype + ) + confidence_features = torch.cat( + [output_hidden_4d, prev_embeddings], + dim=-1, + ) + confidence_pred = self.confidence_head(confidence_features).float() + else: + confidence_pred = self.confidence_head(output_hidden_4d).float() + + return DSparkForwardOutput( + draft_logits=draft_logits, + target_ids=target_ids, + eval_mask=eval_mask, + block_keep_mask=block_keep_mask, + confidence_pred=confidence_pred, + aligned_target_logits=aligned_target_logits, + ) + + +__all__ = [ + "Ministral3DSparkModel", + "Ministral3DSparkAttention", + "Ministral3DSparkDecoderLayer", +] diff --git a/deepspec/modeling/eagle3/__init__.py b/deepspec/modeling/eagle3/__init__.py index a4806ae9..229e51ba 100644 --- a/deepspec/modeling/eagle3/__init__.py +++ b/deepspec/modeling/eagle3/__init__.py @@ -1,10 +1,12 @@ from .common import Eagle3ForwardOutput, extract_eagle3_context_feature from .gemma4 import Gemma4Eagle3Model +from .ministral3 import Ministral3Eagle3Model from .qwen3 import Qwen3Eagle3Model __all__ = [ "Eagle3ForwardOutput", "Gemma4Eagle3Model", + "Ministral3Eagle3Model", "Qwen3Eagle3Model", "extract_eagle3_context_feature", ] diff --git a/deepspec/modeling/eagle3/loss.py b/deepspec/modeling/eagle3/loss.py index 163081d1..1a554fae 100644 --- a/deepspec/modeling/eagle3/loss.py +++ b/deepspec/modeling/eagle3/loss.py @@ -158,6 +158,31 @@ def _log_eagle3_prefix_metrics( add_metric("tau_probabilistic", tau_prob_sum, den=tau_count, tag="train") +def _compute_tv_acceptance_mask( + *, + draft_logits: torch.Tensor, + target_probs: torch.Tensor, + position_mask: torch.Tensor, +) -> torch.Tensor: + draft_probs = torch.softmax(draft_logits.float(), dim=-1) + overlap = torch.minimum(draft_probs, target_probs.float()).sum(dim=-1) + valid_mask = position_mask.squeeze(-1).to(torch.float32) + return overlap * valid_mask + + +def _compute_e2e_tv_loss( + *, + accept_rate_masks: list[torch.Tensor], + start_valid_mask: torch.Tensor, +) -> torch.Tensor: + accept_rate_tensor = torch.stack(accept_rate_masks, dim=0).to(torch.float32) + prefix_acceptance = accept_rate_tensor.cumprod(dim=0).sum(dim=0) + normalized_acceptance = prefix_acceptance / float(len(accept_rate_masks)) + per_position_loss = (1.0 - normalized_acceptance) * start_valid_mask + denominator = start_valid_mask.sum().clamp_min(1.0) + return per_position_loss.sum() / denominator + + @triton.jit def _log_softmax_forward_kernel( logits_ptr, @@ -357,6 +382,7 @@ def compute_eagle3_loss( batch: dict[str, torch.Tensor], ttt_length: int, step_loss_decay: float, + loss_type: str = "soft_ce", ) -> torch.Tensor: input_ids = batch["input_ids"].long() attention_mask = batch["attention_mask"].long() @@ -398,7 +424,13 @@ def compute_eagle3_loss( correct_masks = [] accept_rate_masks = [] + tv_accept_rate_masks = [] valid_masks = [] + loss_type = str(loss_type) + assert loss_type in ("soft_ce", "e2e_tv"), ( + "loss_type must be one of {'soft_ce', 'e2e_tv'}, " + f"got {loss_type!r}." + ) for step_idx in range(int(ttt_length)): # Keep this slice alignment in sync with the Eagle3 reference. target_step_probs = target_probs[ @@ -427,21 +459,43 @@ def compute_eagle3_loss( correct_masks.append(correct_mask) accept_rate_masks.append(accept_rate_mask) valid_masks.append(valid_mask) - step_loss = FusedLogSoftmaxLoss.apply( - output.draft_logits, - target_step_probs, - position_mask_step, - loss_normalizers[step_idx], + if loss_type == "soft_ce": + step_loss = FusedLogSoftmaxLoss.apply( + output.draft_logits, + target_step_probs, + position_mask_step, + loss_normalizers[step_idx], + ) + add_metric( + f"ploss_{step_idx}", + step_loss.detach(), + reduction="dp_mean", + tag="train", + ) + step_weight = float(step_loss_decay) ** step_idx + total_loss = total_loss + step_loss * step_weight + else: + tv_accept_rate_masks.append( + _compute_tv_acceptance_mask( + draft_logits=output.draft_logits, + target_probs=target_step_probs, + position_mask=position_mask_step, + ) + ) + current_input_ids = _shift_with_zero_padding(current_input_ids, left=False) + + if loss_type == "e2e_tv": + start_valid_mask = valid_masks[0].to(torch.float32) + total_loss = _compute_e2e_tv_loss( + accept_rate_masks=tv_accept_rate_masks, + start_valid_mask=start_valid_mask, ) add_metric( - f"ploss_{step_idx}", - step_loss.detach(), + "e2e_tv_loss", + total_loss.detach(), reduction="dp_mean", tag="train", ) - step_weight = float(step_loss_decay) ** step_idx - total_loss = total_loss + step_loss * step_weight - current_input_ids = _shift_with_zero_padding(current_input_ids, left=False) _log_eagle3_prefix_metrics( correct_masks=correct_masks, diff --git a/deepspec/modeling/eagle3/ministral3/__init__.py b/deepspec/modeling/eagle3/ministral3/__init__.py new file mode 100644 index 00000000..0724c5b4 --- /dev/null +++ b/deepspec/modeling/eagle3/ministral3/__init__.py @@ -0,0 +1,7 @@ +from .config import build_draft_config +from .modeling import Ministral3Eagle3Model + +__all__ = [ + "Ministral3Eagle3Model", + "build_draft_config", +] diff --git a/deepspec/modeling/eagle3/ministral3/config.py b/deepspec/modeling/eagle3/ministral3/config.py new file mode 100644 index 00000000..16c445b6 --- /dev/null +++ b/deepspec/modeling/eagle3/ministral3/config.py @@ -0,0 +1,55 @@ +import copy + +from deepspec.modeling.eagle3.common import validate_eagle3_target_layer_ids + + +TRAIN_ATTN_IMPLEMENTATION = "flex_attention" + + +def build_draft_config(*, target_config, model_args): + assert target_config.model_type == "mistral3", ( + "Ministral3 Eagle3 expects the top-level Mistral3 target config, " + f"got model_type={target_config.model_type!r}." + ) + text_config = copy.deepcopy(target_config.text_config) + assert text_config.model_type == "ministral3", ( + "Ministral3 Eagle3 expects target_config.text_config.model_type to be " + f"'ministral3', got {text_config.model_type!r}." + ) + target_layer_ids = validate_eagle3_target_layer_ids( + model_args.target_layer_ids, + int(text_config.num_hidden_layers), + ) + ttt_length = int(model_args.ttt_length) + assert ttt_length >= 1, f"ttt_length must be >= 1, got {ttt_length}" + step_loss_decay = float(model_args.step_loss_decay) + assert step_loss_decay > 0.0, ( + "step_loss_decay must be > 0.0, " + f"got {step_loss_decay}" + ) + draft_num_hidden_layers = int(model_args.draft_num_hidden_layers) + assert draft_num_hidden_layers >= 1, ( + "draft_num_hidden_layers must be >= 1, " + f"got {draft_num_hidden_layers}" + ) + + draft_config = text_config + draft_config.architectures = ["Ministral3Eagle3Model"] + draft_config.target_model_type = str(target_config.model_type) + draft_config.target_text_model_type = str(text_config.model_type) + draft_config.num_target_layers = int(text_config.num_hidden_layers) + draft_config.num_hidden_layers = draft_num_hidden_layers + draft_config.layer_types = ["full_attention"] * draft_num_hidden_layers + draft_config.target_model_name_or_path = str(model_args.target_model_name_or_path) + draft_config.target_layer_ids = target_layer_ids + draft_config.ttt_length = ttt_length + draft_config.step_loss_decay = step_loss_decay + draft_config.draft_num_hidden_layers = draft_num_hidden_layers + draft_config.tie_word_embeddings = False + draft_config._attn_implementation = TRAIN_ATTN_IMPLEMENTATION + return draft_config + + +__all__ = [ + "build_draft_config", +] diff --git a/deepspec/modeling/eagle3/ministral3/modeling.py b/deepspec/modeling/eagle3/ministral3/modeling.py new file mode 100644 index 00000000..c4f416c3 --- /dev/null +++ b/deepspec/modeling/eagle3/ministral3/modeling.py @@ -0,0 +1,423 @@ +from __future__ import annotations + +from typing import Callable, Optional + +import torch +from torch import nn +from torch.nn.attention.flex_attention import flex_attention + +from transformers.cache_utils import Cache +from transformers.models.ministral3.modeling_ministral3 import ( + ALL_ATTENTION_FUNCTIONS, + FlashAttentionKwargs, + GradientCheckpointingLayer, + Ministral3MLP, + Ministral3PreTrainedModel, + Ministral3RMSNorm, + Ministral3RotaryEmbedding, + eager_attention_forward, + get_llama_4_attn_scale, + rotate_half, +) +from typing_extensions import Tuple, Unpack + +from deepspec.modeling.eagle3.common import ( + Eagle3ForwardOutput, + compile_friendly_flex_attention, + create_eagle3_attention_mask, + eagle3_prepare_position_ids, + prepare_4d_causal_attention_mask, +) +from deepspec.utils.sampling import sample_tokens + + +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + q_embed = (q * cos) + (rotate_half(q) * sin) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed + + +class Ministral3Eagle3Attention(nn.Module): + def __init__(self, config, layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = int(layer_idx) + self.head_dim = getattr( + config, + "head_dim", + config.hidden_size // config.num_attention_heads, + ) + self.num_key_value_groups = ( + config.num_attention_heads // config.num_key_value_heads + ) + self.scaling = self.head_dim**-0.5 + self.attention_dropout = config.attention_dropout + self.is_causal = False + + input_dim = int(config.hidden_size) * 2 + self.q_proj = nn.Linear( + input_dim, + config.num_attention_heads * self.head_dim, + bias=False, + ) + self.k_proj = nn.Linear( + input_dim, + config.num_key_value_heads * self.head_dim, + bias=False, + ) + self.v_proj = nn.Linear( + input_dim, + config.num_key_value_heads * self.head_dim, + bias=False, + ) + self.o_proj = nn.Linear( + config.num_attention_heads * self.head_dim, + config.hidden_size, + bias=False, + ) + self.sliding_window = None + + def forward( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: Optional[torch.Tensor], + position_ids: torch.LongTensor, + past_key_values: Optional[Cache] = None, + cache_position: Optional[torch.LongTensor] = None, + past_seen_tokens: int = 0, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + del cache_position + bsz, q_len = hidden_states.shape[:-1] + q = self.q_proj(hidden_states).view(bsz, q_len, -1, self.head_dim) + k = self.k_proj(hidden_states).view(bsz, q_len, -1, self.head_dim) + v = self.v_proj(hidden_states).view(bsz, q_len, -1, self.head_dim) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + + cos, sin = position_embeddings + q, k = apply_rotary_pos_emb(q, k, cos, sin) + q = q * get_llama_4_attn_scale( + position_ids, + self.config.rope_parameters.get("llama_4_scaling_beta"), + self.config.rope_parameters.get("original_max_position_embeddings"), + ).to(q.dtype) + if past_key_values is not None: + k, v = past_key_values.update(k, v, self.layer_idx) + + if self.config._attn_implementation == "flex_attention": + # Direct flex_attention dispatch follows + # SpecForge/specforge/modeling/draft/llama3_eagle.py. + assert attention_mask is not None, ( + "Eagle3 flex_attention expects a BlockMask attention_mask." + ) + flex_attention_func = ( + flex_attention + if int(q_len) <= 128 + else compile_friendly_flex_attention + ) + attn_output = flex_attention_func( + query=q, + key=k.contiguous(), + value=v.contiguous(), + block_mask=attention_mask, + enable_gqa=True, + ) + attn_output = attn_output.transpose(1, 2).contiguous() + attn_output = attn_output.reshape(bsz, q_len, -1) + return self.o_proj(attn_output), None + + attn_fn: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, + eager_attention_forward, + ) + + attn_is_causal = bool( + kwargs.get( + "is_causal", + attention_mask is None and q_len > 1 and int(past_seen_tokens) == 0, + ) + ) + self.is_causal = attn_is_causal + kwargs["is_causal"] = attn_is_causal + attn_output, attn_weights = attn_fn( + self, + q, + k, + v, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + sliding_window=self.sliding_window, + **kwargs, + ) + attn_output = attn_output.reshape(bsz, q_len, -1) + return self.o_proj(attn_output), attn_weights + + +class Ministral3Eagle3DecoderLayer(GradientCheckpointingLayer): + def __init__(self, config, layer_idx: int): + super().__init__() + self.hidden_size = config.hidden_size + self.self_attn = Ministral3Eagle3Attention(config=config, layer_idx=layer_idx) + self.mlp = Ministral3MLP(config) + self.hidden_norm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.input_layernorm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = Ministral3RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + + def forward( + self, + input_embeds: torch.Tensor, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_value: Optional[Cache] = None, + output_attentions: Optional[bool] = False, + use_cache: Optional[bool] = False, + cache_position: Optional[torch.LongTensor] = None, + position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + past_seen_tokens: int = 0, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> torch.Tensor: + del output_attentions, use_cache + assert position_ids is not None, "position_ids must be provided." + residual = hidden_states + hidden_states = self.hidden_norm(hidden_states) + input_embeds = self.input_layernorm(input_embeds) + hidden_states = torch.cat((input_embeds, hidden_states), dim=-1) + hidden_states = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_value, + cache_position=cache_position, + position_embeddings=position_embeddings, + past_seen_tokens=past_seen_tokens, + **kwargs, + )[0] + hidden_states = residual + hidden_states + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + return residual + hidden_states + + +class Ministral3Eagle3Model(Ministral3PreTrainedModel): + # Architecture adapted from + # SpecForge/specforge/modeling/draft/llama3_eagle.py and + # SpecForge/eval/model/eagle3.py. + _no_split_modules = ["Ministral3Eagle3DecoderLayer"] + + def __init__(self, config) -> None: + super().__init__(config) + self.config = config + required_fields = ( + "target_layer_ids", + "ttt_length", + "step_loss_decay", + ) + for field in required_fields: + assert hasattr(config, field), f"config.{field} must be provided." + self.target_layer_ids = [int(x) for x in config.target_layer_ids] + self.ttt_length = int(config.ttt_length) + self.step_loss_decay = float(config.step_loss_decay) + + self.embed_tokens = nn.Embedding( + config.vocab_size, + config.hidden_size, + padding_idx=getattr(config, "pad_token_id", None), + ) + self.fc = nn.Linear( + len(self.target_layer_ids) * config.hidden_size, + config.hidden_size, + bias=False, + ) + self.layers = nn.ModuleList( + [ + Ministral3Eagle3DecoderLayer(config, layer_idx) + for layer_idx in range(config.num_hidden_layers) + ] + ) + self.norm = Ministral3RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.rotary_emb = Ministral3RotaryEmbedding(config) + self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) + self.post_init() + + def initialize_embeddings_and_head( + self, + *, + embed_tokens: nn.Module, + lm_head: nn.Module, + freeze: bool = True, + ): + assert self.embed_tokens.weight.shape == embed_tokens.weight.shape + assert self.lm_head.weight.shape == lm_head.weight.shape + with torch.no_grad(): + self.embed_tokens.weight.copy_(embed_tokens.weight.detach()) + self.lm_head.weight.copy_(lm_head.weight.detach()) + if freeze: + self.set_embedding_head_trainable(False) + + def set_embedding_head_trainable(self, trainable: bool): + self.embed_tokens.requires_grad_(trainable) + self.lm_head.requires_grad_(trainable) + + def project_hidden_states(self, hidden_states: torch.Tensor) -> torch.Tensor: + assert hidden_states.size(-1) == len(self.target_layer_ids) * self.config.hidden_size + return self.fc(hidden_states) + + def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.lm_head(self.norm(hidden_states)) + + def draft_sample(self, logits: torch.Tensor, temperature: float = 0.0) -> torch.Tensor: + return sample_tokens(logits, temperature=temperature) + + def _prepare_attention_mask( + self, + *, + attention_mask: Optional[torch.Tensor], + hidden_states: torch.Tensor, + q_len: int, + past_seen_tokens: int, + ): + if attention_mask is None or attention_mask.ndim != 2: + return attention_mask + kv_len = int(past_seen_tokens) + int(q_len) + if self.config._attn_implementation == "flex_attention": + assert int(past_seen_tokens) % int(q_len) == 0, ( + "Eagle3 flex_attention expects fixed-size TTT chunks: " + f"past_seen_tokens={past_seen_tokens}, q_len={q_len}" + ) + lck = int(past_seen_tokens) // int(q_len) + return create_eagle3_attention_mask( + attention_mask=attention_mask, + q_len=q_len, + kv_len=kv_len, + lck=lck, + device=hidden_states.device, + ) + return prepare_4d_causal_attention_mask( + attention_mask=attention_mask, + dtype=hidden_states.dtype, + q_len=q_len, + kv_len=kv_len, + past_seen_tokens=int(past_seen_tokens), + device=hidden_states.device, + ) + + def extend_draft_cache( + self, + hidden_states: torch.Tensor, + input_ids: torch.LongTensor, + position_ids: torch.LongTensor, + past_key_values: Cache, + ) -> torch.Tensor: + assert input_ids.shape[1] > 0, "input_ids must contain at least one token." + output = self( + hidden_states=hidden_states, + input_ids=input_ids, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=True, + ) + return output[:, -1:, :] + + def forward( + self, + hidden_states: Optional[torch.Tensor] = None, + input_ids: Optional[torch.LongTensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_values: Optional[Cache] = None, + use_cache: bool = False, + attention_mask: Optional[torch.Tensor] = None, + input_embeds: Optional[torch.Tensor] = None, + target_last_hidden_states: Optional[torch.Tensor] = None, + target_logits_only: bool = False, + return_logits: bool = False, + rope_cache_step_offset: bool = False, + **kwargs, + ) -> torch.Tensor | Eagle3ForwardOutput: + if target_logits_only: + assert target_last_hidden_states is not None + with torch.no_grad(): + return self.lm_head(target_last_hidden_states) + + assert hidden_states is not None, "hidden_states must be provided." + if hidden_states.size(-1) == len(self.target_layer_ids) * self.config.hidden_size: + hidden_states = self.project_hidden_states(hidden_states) + + if input_embeds is None: + assert input_ids is not None, "Either input_ids or input_embeds must be provided." + input_embeds = self.embed_tokens(input_ids) + if position_ids is None: + position_ids = eagle3_prepare_position_ids( + input_ids=input_ids, + input_embeds=input_embeds, + ) + + q_len = int(hidden_states.shape[1]) + past_seen_tokens = ( + int(past_key_values.get_seq_length()) + if past_key_values is not None + else 0 + ) + cache_position = torch.arange( + past_seen_tokens, + past_seen_tokens + q_len, + device=hidden_states.device, + ) + prepared_attention_mask = self._prepare_attention_mask( + attention_mask=attention_mask, + hidden_states=hidden_states, + q_len=q_len, + past_seen_tokens=past_seen_tokens, + ) + rope_position_ids = position_ids + if rope_cache_step_offset: + assert int(past_seen_tokens) % int(q_len) == 0, ( + "SpecForge-style Eagle3 RoPE offset expects fixed-size TTT chunks: " + f"past_seen_tokens={past_seen_tokens}, q_len={q_len}" + ) + rope_position_ids = position_ids + int(past_seen_tokens) // int(q_len) + position_embeddings = self.rotary_emb(hidden_states, rope_position_ids) + + for layer in self.layers: + hidden_states = layer( + input_embeds=input_embeds, + hidden_states=hidden_states, + attention_mask=prepared_attention_mask, + position_ids=position_ids, + past_key_value=past_key_values, + use_cache=use_cache, + cache_position=cache_position, + position_embeddings=position_embeddings, + past_seen_tokens=past_seen_tokens, + **kwargs, + ) + if return_logits: + draft_logits = self.compute_logits(hidden_states) + target_logits = None + if target_last_hidden_states is not None: + with torch.no_grad(): + target_logits = self.lm_head(target_last_hidden_states) + return Eagle3ForwardOutput( + hidden_states=hidden_states, + draft_logits=draft_logits, + target_logits=target_logits, + ) + return hidden_states + + +__all__ = [ + "Ministral3Eagle3Model", + "Ministral3Eagle3Attention", + "Ministral3Eagle3DecoderLayer", + "apply_rotary_pos_emb", +] diff --git a/deepspec/modeling/target_utils.py b/deepspec/modeling/target_utils.py new file mode 100644 index 00000000..b8ffe1e5 --- /dev/null +++ b/deepspec/modeling/target_utils.py @@ -0,0 +1,69 @@ +from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer +from transformers.models.mistral3.modeling_mistral3 import ( + Mistral3ForConditionalGeneration, +) + + +def is_mistral3_target(target_model_or_config) -> bool: + if hasattr(target_model_or_config, "config"): + model_type = target_model_or_config.config.model_type + else: + model_type = target_model_or_config.model_type + return str(model_type) == "mistral3" + + +def get_target_backbone(target_model): + model_type = str(target_model.config.model_type) + if model_type in ("gemma4", "gemma4_unified"): + if hasattr(target_model, "language_model"): + return target_model.language_model + if hasattr(target_model, "model") and hasattr( + target_model.model, + "language_model", + ): + return target_model.model.language_model + assert False, "Gemma4 target model must expose a text language_model." + if model_type == "mistral3": + if hasattr(target_model, "language_model"): + return target_model.language_model + if hasattr(target_model, "model") and hasattr( + target_model.model, + "language_model", + ): + return target_model.model.language_model + assert False, "Mistral3 target model must expose a text language_model." + return getattr(target_model, "model", target_model) + + +def get_target_hidden_size(target_model) -> int: + model_type = str(target_model.config.model_type) + if model_type in ("gemma4", "gemma4_unified", "mistral3"): + return int(target_model.config.text_config.hidden_size) + return int(target_model.config.hidden_size) + + +def get_target_input_embeddings(target_model): + if str(target_model.config.model_type) == "mistral3": + return get_target_backbone(target_model).embed_tokens + return target_model.get_input_embeddings() + + +def get_target_output_embeddings(target_model): + return target_model.get_output_embeddings() + + +def load_target_causal_lm(model_name_or_path: str, **kwargs): + target_config = AutoConfig.from_pretrained(model_name_or_path) + if str(target_config.model_type) == "mistral3": + return Mistral3ForConditionalGeneration.from_pretrained( + model_name_or_path, + **kwargs, + ) + return AutoModelForCausalLM.from_pretrained(model_name_or_path, **kwargs) + + +def load_target_tokenizer(model_name_or_path: str, **kwargs): + target_config = AutoConfig.from_pretrained(model_name_or_path) + if str(target_config.model_type) == "mistral3": + kwargs.setdefault("fix_mistral_regex", True) + return AutoTokenizer.from_pretrained(model_name_or_path, **kwargs) diff --git a/deepspec/trainer/__init__.py b/deepspec/trainer/__init__.py index 5ed28b03..0f8ccbbf 100644 --- a/deepspec/trainer/__init__.py +++ b/deepspec/trainer/__init__.py @@ -1,11 +1,21 @@ from .base_trainer import BaseTrainer -from .dspark_trainer import Gemma4DSparkTrainer, Qwen3DSparkTrainer -from .eagle3_trainer import Gemma4Eagle3Trainer, Qwen3Eagle3Trainer +from .dspark_trainer import ( + Gemma4DSparkTrainer, + Ministral3DSparkTrainer, + Qwen3DSparkTrainer, +) +from .eagle3_trainer import ( + Gemma4Eagle3Trainer, + Ministral3Eagle3Trainer, + Qwen3Eagle3Trainer, +) __all__ = [ "BaseTrainer", "Gemma4Eagle3Trainer", "Gemma4DSparkTrainer", + "Ministral3DSparkTrainer", + "Ministral3Eagle3Trainer", "Qwen3Eagle3Trainer", "Qwen3DSparkTrainer", ] diff --git a/deepspec/trainer/base_trainer.py b/deepspec/trainer/base_trainer.py index 6b2eef1d..d4602e4f 100644 --- a/deepspec/trainer/base_trainer.py +++ b/deepspec/trainer/base_trainer.py @@ -1,6 +1,7 @@ from contextlib import nullcontext import math import os +import shutil import torch import torch.distributed as dist @@ -159,7 +160,8 @@ def __init__(self, local_rank, args): self.suspend_controller = SuspendController(device=self.device) self.next_micro_step = 0 - if is_global_main_process(): ensure_dir(self.checkpoint_dir_root) + if is_global_main_process(): + ensure_dir(self.checkpoint_dir_root) training_logger.init( logging_steps=int(self.args.logging.logging_steps), tensorboard_dir=self.args.logging.tensorboard_dir, @@ -174,6 +176,28 @@ def __init__(self, local_rank, args): precision_dtype=self.precision_dtype, global_rank=self.global_rank, ) + else: + init_draft_model_name_or_path = getattr( + self.args.model, + "init_draft_model_name_or_path", + None, + ) + if init_draft_model_name_or_path is not None: + print_on_local_main( + f"Initializing draft model from {init_draft_model_name_or_path}." + ) + self.draft_model = type(self.draft_model).from_pretrained( + str(init_draft_model_name_or_path), + dtype=self.precision_dtype, + attn_implementation=str( + self.draft_model.config._attn_implementation + ), + ) + self.draft_model = self.draft_model.to( + device=self.device, + dtype=self.precision_dtype, + ) + self.draft_model.set_embedding_head_trainable(False) self.model = self.draft_model if self.args.train.torch_compile: print_on_local_main("Compiling training model with torch.compile...") @@ -294,16 +318,22 @@ def _build_train_dataloader(self, start_offset_samples=0, num_samples=None): start_global_offset_samples=start_offset_samples, num_samples=num_samples, ) + num_workers = int(self.args.data.num_workers) + worker_kwargs = {} + if num_workers > 0: + worker_kwargs = { + "persistent_workers": True, + "prefetch_factor": 4, + } return DataLoader( self.train_dataset, batch_size=int(self.args.train.local_batch_size), sampler=sampler, collate_fn=self.data_collator_cls(), - num_workers=int(self.args.data.num_workers), + num_workers=num_workers, pin_memory=True, drop_last=True, - persistent_workers=True, - prefetch_factor=4, + **worker_kwargs, ) def run_batch(self, batch): @@ -323,8 +353,46 @@ def _checkpoint_kwargs(self): local_batch_size=int(self.args.train.local_batch_size), ) - def save_and_eval_checkpoint(self): + def save_train_checkpoint(self): checkpoint_dir = save_checkpoint(**self._checkpoint_kwargs()) + self._prune_checkpoints() + dist.barrier() + return checkpoint_dir + + def _prune_checkpoints(self): + keep_last_checkpoints = int( + getattr(self.args.logging, "keep_last_checkpoints", 0) or 0 + ) + if keep_last_checkpoints <= 0: + return + if not is_global_main_process(): + return + + checkpoints = [] + for name in os.listdir(self.checkpoint_dir_root): + if not name.startswith("step_"): + continue + checkpoint_path = os.path.join(self.checkpoint_dir_root, name) + if not os.path.isdir(checkpoint_path) or os.path.islink(checkpoint_path): + continue + try: + step = int(name.removeprefix("step_")) + except ValueError: + continue + checkpoints.append((step, checkpoint_path)) + + checkpoints.sort() + latest_checkpoint_path = os.path.realpath( + os.path.join(self.checkpoint_dir_root, "step_latest") + ) + for _, checkpoint_path in checkpoints[:-keep_last_checkpoints]: + if os.path.realpath(checkpoint_path) == latest_checkpoint_path: + continue + shutil.rmtree(checkpoint_path) + print_on_global_main(f"Pruned old checkpoint {checkpoint_path}") + + def save_and_eval_checkpoint(self): + checkpoint_dir = self.save_train_checkpoint() if is_global_main_process(): _launch_eval( target_model_name_or_path=self.args.model.target_model_name_or_path, @@ -338,8 +406,7 @@ def save_and_eval_checkpoint(self): def _save_and_suspend(self): print_on_global_main("Saving checkpoint before suspending...") - save_checkpoint(**self._checkpoint_kwargs()) - dist.barrier() + self.save_train_checkpoint() if is_global_main_process(): print_on_global_main("Going to suspend...") self.suspend_controller.go_suspend() @@ -390,8 +457,20 @@ def train(self): grad_norm=grad_norm.item(), ) - if self.global_step % int(self.args.logging.checkpointing_steps) == 0: + should_eval_checkpoint = ( + self.global_step % int(self.args.logging.checkpointing_steps) == 0 + ) + save_only_checkpointing_steps = int( + getattr(self.args.logging, "save_only_checkpointing_steps", 0) or 0 + ) + should_save_only_checkpoint = ( + save_only_checkpointing_steps > 0 + and self.global_step % save_only_checkpointing_steps == 0 + ) + if should_eval_checkpoint: self.save_and_eval_checkpoint() + elif should_save_only_checkpoint: + self.save_train_checkpoint() if self.suspend_controller.requested(): self._save_and_suspend() diff --git a/deepspec/trainer/dspark_trainer.py b/deepspec/trainer/dspark_trainer.py index 487e99a2..9a7d2347 100644 --- a/deepspec/trainer/dspark_trainer.py +++ b/deepspec/trainer/dspark_trainer.py @@ -1,13 +1,25 @@ +from transformers import AutoConfig + from deepspec.data import CacheCollator from deepspec.modeling.dspark.gemma4 import Gemma4DSparkModel from deepspec.modeling.dspark.gemma4.config import ( build_draft_config as build_gemma4_draft_config, ) from deepspec.modeling.dspark.loss import compute_dspark_loss +from deepspec.modeling.dspark.ministral3 import Ministral3DSparkModel +from deepspec.modeling.dspark.ministral3.config import ( + build_draft_config as build_ministral3_draft_config, +) from deepspec.modeling.dspark.qwen3 import Qwen3DSparkModel from deepspec.modeling.dspark.qwen3.config import ( build_draft_config as build_qwen3_draft_config, ) +from deepspec.modeling.target_utils import ( + get_target_input_embeddings, + get_target_output_embeddings, + load_target_causal_lm, + load_target_tokenizer, +) from deepspec.trainer.base_trainer import BaseTrainer @@ -46,3 +58,44 @@ def _build_draft_model(self, *, target_config, model_args): model_args=model_args, ) return Gemma4DSparkModel(draft_config) + + +class Ministral3DSparkTrainer(Qwen3DSparkTrainer): + def build_models(self): + model_args = self.args.model + + tokenizer = load_target_tokenizer( + model_args.target_model_name_or_path, + ) + target_config = AutoConfig.from_pretrained( + model_args.target_model_name_or_path, + ) + + draft_model = self._build_draft_model( + target_config=target_config, + model_args=model_args, + ) + draft_model = draft_model.to(device=self.device, dtype=self.precision_dtype) + + target_model = load_target_causal_lm( + model_args.target_model_name_or_path, + dtype=self.precision_dtype, + ).to(device="cpu").eval() + target_embed_tokens = get_target_input_embeddings(target_model) + target_lm_head = get_target_output_embeddings(target_model) + assert (target_lm_head is not None) and (target_embed_tokens is not None) + draft_model.initialize_embeddings_and_head( + embed_tokens=target_embed_tokens, + lm_head=target_lm_head, + freeze=True, + ) + + del target_model + return draft_model, tokenizer + + def _build_draft_model(self, *, target_config, model_args): + draft_config = build_ministral3_draft_config( + target_config=target_config, + model_args=model_args, + ) + return Ministral3DSparkModel(draft_config) diff --git a/deepspec/trainer/eagle3_trainer.py b/deepspec/trainer/eagle3_trainer.py index 70343a3b..6afe3659 100644 --- a/deepspec/trainer/eagle3_trainer.py +++ b/deepspec/trainer/eagle3_trainer.py @@ -1,4 +1,4 @@ -from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer +from transformers import AutoConfig from deepspec.data import CacheCollator from deepspec.modeling.eagle3.gemma4 import Gemma4Eagle3Model @@ -6,10 +6,20 @@ build_draft_config as build_gemma4_eagle3_config, ) from deepspec.modeling.eagle3.loss import compute_eagle3_loss +from deepspec.modeling.eagle3.ministral3 import Ministral3Eagle3Model +from deepspec.modeling.eagle3.ministral3.config import ( + build_draft_config as build_ministral3_eagle3_config, +) from deepspec.modeling.eagle3.qwen3 import Qwen3Eagle3Model from deepspec.modeling.eagle3.qwen3.config import ( build_draft_config as build_qwen3_eagle3_config, ) +from deepspec.modeling.target_utils import ( + get_target_input_embeddings, + get_target_output_embeddings, + load_target_causal_lm, + load_target_tokenizer, +) from deepspec.trainer.base_trainer import BaseTrainer @@ -19,7 +29,7 @@ class Qwen3Eagle3Trainer(BaseTrainer): def build_models(self): model_args = self.args.model - tokenizer = AutoTokenizer.from_pretrained( + tokenizer = load_target_tokenizer( model_args.target_model_name_or_path, ) target_config = AutoConfig.from_pretrained( @@ -32,12 +42,12 @@ def build_models(self): ) draft_model = draft_model.to(device=self.device, dtype=self.precision_dtype) - target_model = AutoModelForCausalLM.from_pretrained( + target_model = load_target_causal_lm( model_args.target_model_name_or_path, dtype=self.precision_dtype, ).to(device="cpu").eval() - target_embed_tokens = target_model.get_input_embeddings() - target_lm_head = target_model.get_output_embeddings() + target_embed_tokens = get_target_input_embeddings(target_model) + target_lm_head = get_target_output_embeddings(target_model) assert (target_lm_head is not None) and (target_embed_tokens is not None) # The draft head and norm stay frozen / target-independent to match @@ -64,6 +74,7 @@ def run_batch(self, batch): batch=batch, ttt_length=int(self.draft_model.ttt_length), step_loss_decay=float(self.draft_model.step_loss_decay), + loss_type=str(getattr(self.args.model, "loss_type", "soft_ce")), ) @@ -74,3 +85,12 @@ def _build_draft_model(self, *, target_config, model_args): model_args=model_args, ) return Gemma4Eagle3Model(draft_config) + + +class Ministral3Eagle3Trainer(Qwen3Eagle3Trainer): + def _build_draft_model(self, *, target_config, model_args): + draft_config = build_ministral3_eagle3_config( + target_config=target_config, + model_args=model_args, + ) + return Ministral3Eagle3Model(draft_config) diff --git a/eval.py b/eval.py index a35e7ea3..9d0e28d4 100644 --- a/eval.py +++ b/eval.py @@ -1,17 +1,31 @@ from __future__ import annotations import argparse import json +import os + +os.environ["USE_HUB_KERNELS"] = "0" + import torch from transformers import AutoConfig -from deepspec.eval.dspark import Gemma4DSparkEvaluator, Qwen3DSparkEvaluator -from deepspec.eval.eagle3 import Gemma4Eagle3Evaluator, Qwen3Eagle3Evaluator +from deepspec.eval.dspark import ( + Gemma4DSparkEvaluator, + Ministral3DSparkEvaluator, + Qwen3DSparkEvaluator, +) +from deepspec.eval.eagle3 import ( + Gemma4Eagle3Evaluator, + Ministral3Eagle3Evaluator, + Qwen3Eagle3Evaluator, +) from deepspec.utils import CustomJSONEncoder EVALUATORS = { "Qwen3DSparkModel": Qwen3DSparkEvaluator, "Gemma4DSparkModel": Gemma4DSparkEvaluator, + "Ministral3DSparkModel": Ministral3DSparkEvaluator, "Qwen3Eagle3Model": Qwen3Eagle3Evaluator, "Gemma4Eagle3Model": Gemma4Eagle3Evaluator, + "Ministral3Eagle3Model": Ministral3Eagle3Evaluator, "Eagle3DraftModel": Qwen3Eagle3Evaluator, } @@ -27,7 +41,18 @@ ("arena-hard-v2", 500), ] -def parse_args(): + +def parse_task_spec(task_spec: str) -> tuple[str, int | None]: + if ":" not in task_spec: + return task_spec, None + + name, max_samples = task_spec.rsplit(":", 1) + if max_samples.lower() == "none": + return name, None + return name, int(max_samples) + + +def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--target_name_or_path", type=str, required=True) parser.add_argument("--draft_name_or_path",type=str,required=True) @@ -40,10 +65,23 @@ def parse_args(): help=("Confidence-head early-stop threshold. Confidence calibration metrics are collected only when this is 0.0."), ) parser.add_argument("--tensorboard-dir", type=str, default=None) + parser.add_argument("--output-json", type=str, default=None) parser.add_argument("--step", type=int, default=None,help=("step for tensorboard logging"),) parser.add_argument("--seed", type=int, default=980406) + parser.add_argument( + "--task", + action="append", + default=None, + help=( + "Evaluate only this task. Repeatable. Format is `name` to use all " + "available samples or `name:max_samples` to cap the dataset." + ), + ) args = parser.parse_args() - args.tasks = list(TASKS) + if args.task is None: + args.tasks = list(TASKS) + else: + args.tasks = [parse_task_spec(task_spec) for task_spec in args.task] return args diff --git a/experiments/ministral3_3b_eagle3/README.md b/experiments/ministral3_3b_eagle3/README.md new file mode 100644 index 00000000..fcf8fcb4 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/README.md @@ -0,0 +1,30 @@ +# Ministral3 EAGLE3 Reproduction + +This directory tracks the DeepSpec-side reproduction attempt for +`mistralai/Ministral-3-3B-Instruct-2512`. + +The smoke scripts are intentionally tiny: + +- `prepare_cache_smoke.sbatch` builds a four-sample target cache. +- `train_smoke.sbatch` trains the EAGLE3 draft for two optimizer steps against + that cache. + +They validate model loading, target hidden-state extraction, tokenizer/loss-mask +parsing, and the EAGLE3 training loop before launching a full Open-PerfectBlend +run. + +The first reproduction pipeline uses a 10k Open-PerfectBlend subset generated +by the Ministral3 target, then trains for 3000 optimizer steps at DeepSpec's +EAGLE3 global batch size of 512: + +- `download_subset_10k.sbatch` streams and normalizes the source subset. +- `generate_subset_10k.sbatch` regenerates assistant turns with Ministral3. +- `merge_subset_10k.sbatch` combines the eight regeneration shards. +- `prepare_cache_10k.sbatch` builds the target hidden-state cache. +- `train_10k_3000steps.sbatch` trains + `eagle3_ttt7_ministral3_3b_10k_3000steps`. +- `eval_10k_3000steps.sbatch` evaluates the resulting `step_latest`. + +`launch_10k_pipeline.sh` submits those jobs with `afterok` dependencies. Outputs +live under `/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3`, checkpoints under +`~/checkpoints/deepspec`, and TensorBoard logs under `~/tensorboard/deepspec`. diff --git a/experiments/ministral3_3b_eagle3/download_subset_10k.sbatch b/experiments/ministral3_3b_eagle3/download_subset_10k.sbatch new file mode 100755 index 00000000..a5d77ee6 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/download_subset_10k.sbatch @@ -0,0 +1,25 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-subset10k +#SBATCH --partition=cpu +#SBATCH --qos=research +#SBATCH --cpus-per-task=4 +#SBATCH --mem=16G +#SBATCH --time=00:30:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +DATA_DIR=${RUN_ROOT}/data_10k + +mkdir -p "${RUN_ROOT}/logs" "${DATA_DIR}" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +python scripts/data/download_streaming_subset.py \ + --sample-size 10000 \ + --output-path "${DATA_DIR}/perfectblend_train_10k.jsonl" \ + --skip-existing diff --git a/experiments/ministral3_3b_eagle3/eval_10k_3000steps.sbatch b/experiments/ministral3_3b_eagle3/eval_10k_3000steps.sbatch new file mode 100755 index 00000000..d9a7bd49 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/eval_10k_3000steps.sbatch @@ -0,0 +1,36 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-eval10k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +DRAFT_DIR=${HOME}/checkpoints/deepspec/eagle3_ttt7_ministral3_3b_10k_3000steps/step_latest +TB_DIR=${HOME}/tensorboard/deepspec/eagle3_ttt7_ministral3_3b_10k_3000steps_eval + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29671} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_10k_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_eagle3/eval_10k_3000steps_task_array.sbatch b/experiments/ministral3_3b_eagle3/eval_10k_3000steps_task_array.sbatch new file mode 100644 index 00000000..c29c84d9 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/eval_10k_3000steps_task_array.sbatch @@ -0,0 +1,55 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-evaltask +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:4 +#SBATCH --cpus-per-task=24 +#SBATCH --mem=192G +#SBATCH --time=03:00:00 +#SBATCH --array=1-8%2 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%A_%a.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%A_%a.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +DRAFT_DIR=${HOME}/checkpoints/deepspec/eagle3_ttt7_ministral3_3b_10k_3000steps/step_latest +TB_ROOT=${HOME}/tensorboard/deepspec/eagle3_ttt7_ministral3_3b_10k_3000steps_eval_tasks + +TASK_NAMES=( + gsm8k + math500 + aime25 + humaneval + mbpp + livecodebench + mt-bench + alpaca + arena-hard-v2 +) +TASK_COUNTS=(500 500 30 164 256 500 80 500 500) + +TASK_NAME=${TASK_NAMES[${SLURM_ARRAY_TASK_ID}]} +TASK_COUNT=${TASK_COUNTS[${SLURM_ARRAY_TASK_ID}]} +SAFE_TASK_NAME=${TASK_NAME//[^a-zA-Z0-9]/_} + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29681} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_ROOT}/${SAFE_TASK_NAME}" \ + --output-json "${RUN_ROOT}/eval_10k_step3000_${SAFE_TASK_NAME}_metrics.json" \ + --step 3000 \ + --task "${TASK_NAME}:${TASK_COUNT}" diff --git a/experiments/ministral3_3b_eagle3/generate_subset_10k.sbatch b/experiments/ministral3_3b_eagle3/generate_subset_10k.sbatch new file mode 100755 index 00000000..b10e4957 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/generate_subset_10k.sbatch @@ -0,0 +1,38 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-gen10k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:1 +#SBATCH --cpus-per-task=12 +#SBATCH --mem=160G +#SBATCH --time=12:00:00 +#SBATCH --array=0-7 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%A_%a.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%A_%a.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +DATA_DIR=${RUN_ROOT}/data_10k +NUM_SHARDS=8 +SHARD_INDEX=${SLURM_ARRAY_TASK_ID} + +mkdir -p "${RUN_ROOT}/logs" "${DATA_DIR}/regen_shards" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export USE_HUB_KERNELS=0 + +python scripts/data/generate_train_data_transformers.py \ + --model mistralai/Ministral-3-3B-Instruct-2512 \ + --input-file-path "${DATA_DIR}/perfectblend_train_10k.jsonl" \ + --output-file-path "${DATA_DIR}/regen_shards/perfectblend_train_10k_regen_shard_${SHARD_INDEX}.jsonl" \ + --error-file-path "${DATA_DIR}/regen_shards/perfectblend_train_10k_regen_shard_${SHARD_INDEX}_error.jsonl" \ + --shard-index "${SHARD_INDEX}" \ + --num-shards "${NUM_SHARDS}" \ + --temperature 0.7 \ + --top-p 0.8 \ + --top-k 20 \ + --max-new-tokens 4096 \ + --resume diff --git a/experiments/ministral3_3b_eagle3/launch_10k_pipeline.sh b/experiments/ministral3_3b_eagle3/launch_10k_pipeline.sh new file mode 100755 index 00000000..67d2cfda --- /dev/null +++ b/experiments/ministral3_3b_eagle3/launch_10k_pipeline.sh @@ -0,0 +1,20 @@ +#!/usr/bin/env bash +set -euo pipefail + +HERE=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd) + +download_id=$(sbatch --parsable "${HERE}/download_subset_10k.sbatch") +generate_id=$(sbatch --parsable --dependency=afterok:"${download_id}" "${HERE}/generate_subset_10k.sbatch") +merge_id=$(sbatch --parsable --dependency=afterok:"${generate_id}" "${HERE}/merge_subset_10k.sbatch") +cache_id=$(sbatch --parsable --dependency=afterok:"${merge_id}" "${HERE}/prepare_cache_10k.sbatch") +train_id=$(sbatch --parsable --dependency=afterok:"${cache_id}" "${HERE}/train_10k_3000steps.sbatch") +eval_id=$(sbatch --parsable --dependency=afterok:"${train_id}" "${HERE}/eval_10k_3000steps.sbatch") + +cat <> "${OUT}" +done + +wc -l "${OUT}" diff --git a/experiments/ministral3_3b_eagle3/plot_training_metrics.py b/experiments/ministral3_3b_eagle3/plot_training_metrics.py new file mode 100644 index 00000000..9fddbbcd --- /dev/null +++ b/experiments/ministral3_3b_eagle3/plot_training_metrics.py @@ -0,0 +1,198 @@ +import argparse +import json +from pathlib import Path + +import matplotlib.pyplot as plt +from tensorboard.backend.event_processing.event_accumulator import EventAccumulator + + +ScalarSeries = list[tuple[int, float]] + +QWEN3_4B_EAGLE3_ACCEPTANCE_LENGTHS = { + "gsm8k": 5.14, + "math500": 4.62, + "aime25": 3.92, + "mbpp": 3.69, + "humaneval": 4.16, + "livecodebench": 3.77, + "mt-bench": 2.39, + "alpaca": 2.26, + "arena-hard-v2": 2.55, +} + +DATASET_LABELS = { + "gsm8k": "GSM8K", + "math500": "MATH", + "aime25": "AIME25", + "mbpp": "MBPP", + "humaneval": "HumanEval", + "livecodebench": "LCB", + "mt-bench": "MT-Bench", + "alpaca": "Alpaca", + "arena-hard-v2": "Arena-Hard", +} + + +def load_scalars(tensorboard_dir: Path) -> dict[str, ScalarSeries]: + accumulator = EventAccumulator(str(tensorboard_dir), size_guidance={"scalars": 0}) + accumulator.Reload() + + scalars = {} + for tag in accumulator.Tags()["scalars"]: + values_by_step = { + int(event.step): float(event.value) + for event in accumulator.Scalars(tag) + } + scalars[tag] = sorted(values_by_step.items()) + return scalars + + +def plot_series( + *, + scalars: dict[str, ScalarSeries], + tags: list[str], + output_path: Path, + title: str, + ylabel: str, +) -> None: + plt.figure(figsize=(8, 4.5)) + for tag in tags: + series = scalars.get(tag, []) + if not series: + continue + steps = [step for step, _ in series] + values = [value for _, value in series] + label = tag.removeprefix("train/") + plt.plot(steps, values, label=label, linewidth=1.8) + + plt.title(title) + plt.xlabel("Optimizer step") + plt.ylabel(ylabel) + plt.grid(True, alpha=0.25) + plt.legend(loc="best", fontsize=8) + plt.tight_layout() + output_path.parent.mkdir(parents=True, exist_ok=True) + plt.savefig(output_path, dpi=180) + plt.close() + + +def write_summary(scalars: dict[str, ScalarSeries], output_path: Path) -> None: + summary = {} + for tag, series in sorted(scalars.items()): + if not series: + continue + step, value = series[-1] + summary[tag] = {"step": step, "value": value} + + output_path.parent.mkdir(parents=True, exist_ok=True) + output_path.write_text( + json.dumps(summary, indent=2, sort_keys=True, ensure_ascii=False) + "\n" + ) + + +def plot_eval_comparison(eval_json_path: Path, output_dir: Path) -> None: + eval_payload = json.loads(eval_json_path.read_text()) + ministral_by_dataset = { + row["dataset"]: float(row["acceptance_length"]) + for row in eval_payload["rows"] + } + datasets = [ + dataset + for dataset in QWEN3_4B_EAGLE3_ACCEPTANCE_LENGTHS + if dataset in ministral_by_dataset + ] + labels = [DATASET_LABELS[dataset] for dataset in datasets] + ministral_values = [ministral_by_dataset[dataset] for dataset in datasets] + qwen_values = [ + QWEN3_4B_EAGLE3_ACCEPTANCE_LENGTHS[dataset] for dataset in datasets + ] + x_positions = list(range(len(datasets))) + bar_width = 0.38 + + plt.figure(figsize=(10, 4.8)) + plt.bar( + [position - bar_width / 2 for position in x_positions], + ministral_values, + width=bar_width, + label="Ministral3-3B Eagle3", + ) + plt.bar( + [position + bar_width / 2 for position in x_positions], + qwen_values, + width=bar_width, + label="Qwen3-4B Eagle3", + ) + plt.xticks(x_positions, labels, rotation=30, ha="right") + plt.ylabel("Accepted length") + plt.title("Final accepted length by benchmark") + plt.grid(True, axis="y", alpha=0.25) + plt.legend(loc="best", fontsize=8) + plt.tight_layout() + output_dir.mkdir(parents=True, exist_ok=True) + plt.savefig(output_dir / "eval_acceptance_comparison.png", dpi=180) + plt.close() + + comparison_rows = [] + for dataset, ministral_value, qwen_value in zip( + datasets, ministral_values, qwen_values, strict=True + ): + comparison_rows.append( + { + "dataset": dataset, + "label": DATASET_LABELS[dataset], + "ministral3_3b_eagle3": ministral_value, + "qwen3_4b_eagle3": qwen_value, + "delta": ministral_value - qwen_value, + } + ) + output_path = output_dir / "eval_acceptance_comparison.json" + output_path.write_text( + json.dumps(comparison_rows, indent=2, sort_keys=True, ensure_ascii=False) + + "\n" + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--tensorboard-dir", required=True, type=Path) + parser.add_argument("--output-dir", required=True, type=Path) + parser.add_argument("--eval-json", type=Path, default=None) + args = parser.parse_args() + + scalars = load_scalars(args.tensorboard_dir) + + plot_series( + scalars=scalars, + tags=["train/loss"], + output_path=args.output_dir / "train_loss.png", + title="Training loss", + ylabel="Loss", + ) + plot_series( + scalars=scalars, + tags=["train/tau_greedy", "train/tau_probabilistic"], + output_path=args.output_dir / "train_tau.png", + title="Expected accepted length proxies", + ylabel="Tokens", + ) + plot_series( + scalars=scalars, + tags=[f"train/accept_rate@{idx}" for idx in range(7)], + output_path=args.output_dir / "train_accept_rates.png", + title="Per-position acceptance rates", + ylabel="Acceptance rate", + ) + plot_series( + scalars=scalars, + tags=[f"train/accuracy@{idx}" for idx in range(7)], + output_path=args.output_dir / "train_accuracies.png", + title="Per-position token accuracies", + ylabel="Accuracy", + ) + write_summary(scalars, args.output_dir / "train_metric_summary.json") + if args.eval_json is not None: + plot_eval_comparison(args.eval_json, args.output_dir) + + +if __name__ == "__main__": + main() diff --git a/experiments/ministral3_3b_eagle3/prepare_cache_10k.sbatch b/experiments/ministral3_3b_eagle3/prepare_cache_10k.sbatch new file mode 100755 index 00000000..8834f77f --- /dev/null +++ b/experiments/ministral3_3b_eagle3/prepare_cache_10k.sbatch @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-cache10k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +DATA_DIR=${RUN_ROOT}/data_10k +CACHE_DIR=${RUN_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29651} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python scripts/data/prepare_target_cache.py \ + --config config/eagle3/eagle3_ministral3_3b.py \ + --train-data-path "${DATA_DIR}/perfectblend_train_10k_regen.jsonl" \ + --output-dir "${CACHE_DIR}" \ + --local-batch-size 1 \ + --num-workers 2 \ + --max-shard-bytes 17179869184 diff --git a/experiments/ministral3_3b_eagle3/prepare_cache_smoke.sbatch b/experiments/ministral3_3b_eagle3/prepare_cache_smoke.sbatch new file mode 100755 index 00000000..85cbf259 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/prepare_cache_smoke.sbatch @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-cache-smoke +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:1 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=80G +#SBATCH --time=00:20:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${RUN_ROOT}/smoke_target_cache + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29631} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python scripts/data/prepare_target_cache.py \ + --config config/eagle3/eagle3_ministral3_3b.py \ + --train-data-path experiments/ministral3_3b_eagle3/smoke_train.jsonl \ + --output-dir "${CACHE_DIR}" \ + --min-loss-tokens 1 \ + --local-batch-size 1 \ + --num-workers 0 \ + --max-shard-bytes 1073741824 diff --git a/experiments/ministral3_3b_eagle3/report.md b/experiments/ministral3_3b_eagle3/report.md new file mode 100644 index 00000000..85d28eed --- /dev/null +++ b/experiments/ministral3_3b_eagle3/report.md @@ -0,0 +1,66 @@ +# Ministral3-3B Eagle3 DeepSpec Reproduction + +This run integrated `mistralai/Ministral-3-3B-Instruct-2512` into the DeepSpec +Eagle3 training and evaluation stack and trained a TTT-7 Eagle3 draft for 3000 +optimizer steps. The final result does not reproduce the Qwen3-4B Eagle3 +accepted-length numbers reported by DeepSpec: the Ministral3-3B macro accepted +length is `2.40`, versus `3.61` for the Qwen3-4B Eagle3 table. + +## Setup + +| Field | Value | +| --- | --- | +| Target model | `mistralai/Ministral-3-3B-Instruct-2512` | +| Draft recipe | Eagle3, TTT length 7 | +| Target layers | `[1, 7, 13, 19, 24]` | +| Training data | 9,907 target-regenerated Open-PerfectBlend samples | +| Max sequence length | 4096 | +| Global batch size | 512 | +| Training steps | 3000 | +| Learning rate | `6e-4` | +| Warmup ratio | `0.04` | +| Grad clipping | `1.0` | +| Final checkpoint | `/mnt/vast/home/andy/checkpoints/deepspec/eagle3_ttt7_ministral3_3b_10k_3000steps/step_latest` | +| Final eval JSON | `/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/eval_10k_step3000_metrics.json` | + +## Training Curves + +The training metrics reached high in-distribution accept-rate proxies by step +3000: `accept_rate@0..6 = 0.783, 0.780, 0.781, 0.779, 0.773, 0.762, 0.741`, +with `tau_greedy = 5.19` and `tau_probabilistic = 4.43`. + +![Training loss](report_assets/train_loss.png) + +![Training tau proxies](report_assets/train_tau.png) + +![Training accept rates](report_assets/train_accept_rates.png) + +![Training accuracies](report_assets/train_accuracies.png) + +## Final Evaluation + +| Dataset | Ministral3-3B Eagle3 accepted length | DeepSpec Qwen3-4B Eagle3 accepted length | Delta | Verify rate | +| --- | ---: | ---: | ---: | ---: | +| GSM8K | 2.84 | 5.14 | -2.30 | 0.357 | +| MATH500 | 3.10 | 4.62 | -1.52 | 0.388 | +| AIME25 | 2.88 | 3.92 | -1.04 | 0.360 | +| MBPP | 2.46 | 3.69 | -1.23 | 0.310 | +| HumanEval | 2.60 | 4.16 | -1.56 | 0.326 | +| LiveCodeBench | 2.27 | 3.77 | -1.50 | 0.283 | +| MT-Bench | 1.89 | 2.39 | -0.50 | 0.237 | +| Alpaca | 1.83 | 2.26 | -0.43 | 0.230 | +| Arena-Hard-v2 | 1.69 | 2.55 | -0.86 | 0.212 | +| **Macro mean** | **2.40** | **3.61** | **-1.22** | **0.300** | + +![Eval accepted length comparison](report_assets/eval_acceptance_comparison.png) + +## Conclusion + +The implementation is functional and the training run converges on the +target-regenerated subset, but the final evaluation is well below the DeepSpec +Qwen3-4B Eagle3 reference. The train/eval gap is large: training accept-rate +metrics finish near `0.74-0.78`, while eval accepted length ranges from `1.69` +to `3.10`. The most likely next reproduction lever is data scale and +distribution, because this run used a 10k regenerated subset rather than the +full target-generated Open-PerfectBlend training corpus used for the published +DeepSpec checkpoints. diff --git a/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.json b/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.json new file mode 100644 index 00000000..774d49e0 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.json @@ -0,0 +1,65 @@ +[ + { + "dataset": "gsm8k", + "delta": -2.296284255273641, + "label": "GSM8K", + "ministral3_3b_eagle3": 2.843715744726359, + "qwen3_4b_eagle3": 5.14 + }, + { + "dataset": "math500", + "delta": -1.5230024237028243, + "label": "MATH", + "ministral3_3b_eagle3": 3.0969975762971758, + "qwen3_4b_eagle3": 4.62 + }, + { + "dataset": "aime25", + "delta": -1.0425132859897475, + "label": "AIME25", + "ministral3_3b_eagle3": 2.8774867140102525, + "qwen3_4b_eagle3": 3.92 + }, + { + "dataset": "mbpp", + "delta": -1.2256578443930595, + "label": "MBPP", + "ministral3_3b_eagle3": 2.4643421556069405, + "qwen3_4b_eagle3": 3.69 + }, + { + "dataset": "humaneval", + "delta": -1.5595768611779173, + "label": "HumanEval", + "ministral3_3b_eagle3": 2.600423138822083, + "qwen3_4b_eagle3": 4.16 + }, + { + "dataset": "livecodebench", + "delta": -1.5047376374946562, + "label": "LCB", + "ministral3_3b_eagle3": 2.265262362505344, + "qwen3_4b_eagle3": 3.77 + }, + { + "dataset": "mt-bench", + "delta": -0.5001109907958854, + "label": "MT-Bench", + "ministral3_3b_eagle3": 1.8898890092041147, + "qwen3_4b_eagle3": 2.39 + }, + { + "dataset": "alpaca", + "delta": -0.42997873936968456, + "label": "Alpaca", + "ministral3_3b_eagle3": 1.8300212606303152, + "qwen3_4b_eagle3": 2.26 + }, + { + "dataset": "arena-hard-v2", + "delta": -0.855662059518115, + "label": "Arena-Hard", + "ministral3_3b_eagle3": 1.6943379404818848, + "qwen3_4b_eagle3": 2.55 + } +] diff --git a/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.png b/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.png new file mode 100644 index 00000000..0a47ea83 Binary files /dev/null and b/experiments/ministral3_3b_eagle3/report_assets/eval_acceptance_comparison.png differ diff --git a/experiments/ministral3_3b_eagle3/report_assets/train_accept_rates.png b/experiments/ministral3_3b_eagle3/report_assets/train_accept_rates.png new file mode 100644 index 00000000..f1057854 Binary files /dev/null and b/experiments/ministral3_3b_eagle3/report_assets/train_accept_rates.png differ diff --git a/experiments/ministral3_3b_eagle3/report_assets/train_accuracies.png b/experiments/ministral3_3b_eagle3/report_assets/train_accuracies.png new file mode 100644 index 00000000..57aeb124 Binary files /dev/null and b/experiments/ministral3_3b_eagle3/report_assets/train_accuracies.png differ diff --git a/experiments/ministral3_3b_eagle3/report_assets/train_loss.png b/experiments/ministral3_3b_eagle3/report_assets/train_loss.png new file mode 100644 index 00000000..03809967 Binary files /dev/null and b/experiments/ministral3_3b_eagle3/report_assets/train_loss.png differ diff --git a/experiments/ministral3_3b_eagle3/report_assets/train_metric_summary.json b/experiments/ministral3_3b_eagle3/report_assets/train_metric_summary.json new file mode 100644 index 00000000..e60f17aa --- /dev/null +++ b/experiments/ministral3_3b_eagle3/report_assets/train_metric_summary.json @@ -0,0 +1,134 @@ +{ + "train/accept_rate@0": { + "step": 3000, + "value": 0.7828433513641357 + }, + "train/accept_rate@1": { + "step": 3000, + "value": 0.7800154685974121 + }, + "train/accept_rate@2": { + "step": 3000, + "value": 0.7808699011802673 + }, + "train/accept_rate@3": { + "step": 3000, + "value": 0.7791018486022949 + }, + "train/accept_rate@4": { + "step": 3000, + "value": 0.7734428644180298 + }, + "train/accept_rate@5": { + "step": 3000, + "value": 0.7620912194252014 + }, + "train/accept_rate@6": { + "step": 3000, + "value": 0.740893542766571 + }, + "train/accuracy@0": { + "step": 3000, + "value": 0.8395382761955261 + }, + "train/accuracy@1": { + "step": 3000, + "value": 0.8423240184783936 + }, + "train/accuracy@2": { + "step": 3000, + "value": 0.848247230052948 + }, + "train/accuracy@3": { + "step": 3000, + "value": 0.851127028465271 + }, + "train/accuracy@4": { + "step": 3000, + "value": 0.8497721552848816 + }, + "train/accuracy@5": { + "step": 3000, + "value": 0.8434146046638489 + }, + "train/accuracy@6": { + "step": 3000, + "value": 0.8280242085456848 + }, + "train/grad_norm": { + "step": 3000, + "value": 0.078125 + }, + "train/loss": { + "step": 3000, + "value": 2.1320221424102783 + }, + "train/lr": { + "step": 3000, + "value": 0.0 + }, + "train/ploss_0": { + "step": 3000, + "value": 0.5350198745727539 + }, + "train/ploss_1": { + "step": 3000, + "value": 0.535610020160675 + }, + "train/ploss_2": { + "step": 3000, + "value": 0.5331065654754639 + }, + "train/ploss_3": { + "step": 3000, + "value": 0.5346406698226929 + }, + "train/ploss_4": { + "step": 3000, + "value": 0.540634036064148 + }, + "train/ploss_5": { + "step": 3000, + "value": 0.552795946598053 + }, + "train/ploss_6": { + "step": 3000, + "value": 0.5760428309440613 + }, + "train/tau_greedy": { + "step": 3000, + "value": 5.192190647125244 + }, + "train/tau_probabilistic": { + "step": 3000, + "value": 4.4316630363464355 + }, + "train/valid_tokens@0": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@1": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@2": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@3": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@4": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@5": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@6": { + "step": 3000, + "value": 8655.0615234375 + } +} diff --git a/experiments/ministral3_3b_eagle3/report_assets/train_tau.png b/experiments/ministral3_3b_eagle3/report_assets/train_tau.png new file mode 100644 index 00000000..f4afcca4 Binary files /dev/null and b/experiments/ministral3_3b_eagle3/report_assets/train_tau.png differ diff --git a/experiments/ministral3_3b_eagle3/smoke_train.jsonl b/experiments/ministral3_3b_eagle3/smoke_train.jsonl new file mode 100644 index 00000000..1d4ad701 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/smoke_train.jsonl @@ -0,0 +1,4 @@ +{"conversations": [{"role": "user", "content": "What is 2 + 2?"}, {"role": "assistant", "content": "4"}]} +{"conversations": [{"role": "user", "content": "Name the capital of France."}, {"role": "assistant", "content": "Paris"}]} +{"conversations": [{"role": "user", "content": "Write a Python function that returns x squared."}, {"role": "assistant", "content": "def square(x):\n return x * x"}]} +{"conversations": [{"role": "user", "content": "Give one short reason water freezes."}, {"role": "assistant", "content": "Water freezes when it loses enough heat for molecules to form solid ice."}]} diff --git a/experiments/ministral3_3b_eagle3/train_10k_3000steps.sbatch b/experiments/ministral3_3b_eagle3/train_10k_3000steps.sbatch new file mode 100755 index 00000000..22359b13 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/train_10k_3000steps.sbatch @@ -0,0 +1,40 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-train10k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=12:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${RUN_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29661} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b.py \ + --opts "exp_name=eagle3_ttt7_ministral3_3b_10k_3000steps" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=10" \ + --opts "logging.keep_last_checkpoints=8" diff --git a/experiments/ministral3_3b_eagle3/train_10k_segment.sbatch b/experiments/ministral3_3b_eagle3/train_10k_segment.sbatch new file mode 100755 index 00000000..3c5a82a6 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/train_10k_segment.sbatch @@ -0,0 +1,41 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-segment +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${RUN_ROOT}/target_cache_10k +MAX_TRAIN_STEPS=${MAX_TRAIN_STEPS:?Set MAX_TRAIN_STEPS for this segment} + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29661} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b.py \ + --opts "exp_name=eagle3_ttt7_ministral3_3b_10k_3000steps" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=${MAX_TRAIN_STEPS}" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=10" \ + --opts "logging.keep_last_checkpoints=8" diff --git a/experiments/ministral3_3b_eagle3/train_smoke.sbatch b/experiments/ministral3_3b_eagle3/train_smoke.sbatch new file mode 100755 index 00000000..6f892de4 --- /dev/null +++ b/experiments/ministral3_3b_eagle3/train_smoke.sbatch @@ -0,0 +1,38 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-train-smoke +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:1 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=80G +#SBATCH --time=00:20:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${RUN_ROOT}/smoke_target_cache + +mkdir -p "${RUN_ROOT}/logs" +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29641} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b.py \ + --opts "exp_name=eagle3_ttt7_ministral3_3b_smoke" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=0" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=1" \ + --opts "train.max_train_steps=2" \ + --opts "logging.checkpointing_steps=1" diff --git a/experiments/ministral3_3b_followup/README.md b/experiments/ministral3_3b_followup/README.md new file mode 100644 index 00000000..f1e961e4 --- /dev/null +++ b/experiments/ministral3_3b_followup/README.md @@ -0,0 +1,82 @@ +# Ministral3 Follow-Up Experiments + +This follow-up tested whether the DeepSpec DSpark recipe transfers cleanly to +`mistralai/Ministral-3-3B-Instruct-2512`, and kept the Eagle3 KL baseline as the +main reference point. + +## Final Status + +The DSpark block-7 run reached `step_3000`, and the final eval artifact is now +complete with all nine benchmark rows: + +- Checkpoint: `/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b_8gpu/step_3000` +- Final eval JSON: `/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_metrics.json` +- Preserved pre-merge 8/9 eval JSON: `/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_metrics_partial_8of9.json` +- Arena-only rescue JSON: `/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_arena_metrics.json` + +The result does **not** reproduce the expected DSpark improvement over the +Eagle3 KL baseline on Ministral3-3B. DSpark finished with stronger train-side +acceptance proxies, but its final eval accepted length is lower than the KL +baseline on every benchmark. + +## Final Eval + +| Dataset | KL accepted length | DSpark accepted length | Delta | DSpark verify rate | +|---|---:|---:|---:|---:| +| gsm8k | 2.844 | 2.593 | -0.251 | 0.325 | +| math500 | 3.097 | 2.962 | -0.135 | 0.371 | +| aime25 | 2.877 | 2.688 | -0.189 | 0.336 | +| humaneval | 2.600 | 2.435 | -0.165 | 0.305 | +| mbpp | 2.464 | 2.331 | -0.133 | 0.292 | +| livecodebench | 2.265 | 2.120 | -0.146 | 0.265 | +| mt-bench | 1.890 | 1.841 | -0.049 | 0.230 | +| alpaca | 1.830 | 1.775 | -0.055 | 0.222 | +| arena-hard-v2 | 1.694 | 1.581 | -0.113 | 0.198 | +| **Macro avg** | **2.396** | **2.258** | **-0.137** | **0.283** | + +![Final accepted length by benchmark](report_assets/eval_acceptance_comparison.png) + +## Training Metrics + +The DSpark training run ended with strong train-side speculative proxies: + +- `train/loss`: `0.4216` +- `train/tau_probabilistic`: `5.8782` +- `train/accept_rate@0`: `0.9089` +- `train/accept_rate@3`: `0.8717` +- `train/accept_rate@6`: `0.8207` + +That train/eval gap is the main remaining discrepancy: the training accept-rate +metrics look healthy, but the official eval accepted lengths remain low. + +![Training loss](report_assets/train_loss.png) + +![Training tau probabilistic](report_assets/train_tau_probabilistic.png) + +![Training accept rates](report_assets/train_accept_rates.png) + +The scalar snapshot used for these plots is in +[report_assets/train_metric_summary.json](report_assets/train_metric_summary.json). + +## Reproduction Notes + +All heavy artifacts were written under `/mnt/vast/runs/andy`. The final +`arena-hard-v2` row came from a one-dataset rescue eval because the first full +eval hit its 2h walltime after writing 8/9 rows. The rescue completed in +`00:40:44`, then `merge_eval_json.py` validated and merged the one-row arena +JSON into the final complete aggregate. + +Useful commands: + +```bash +sbatch experiments/ministral3_3b_followup/train_dspark_3000steps_8gpu.sbatch +sbatch experiments/ministral3_3b_followup/eval_dspark_8gpu_3000steps.sbatch +sbatch experiments/ministral3_3b_followup/eval_dspark_8gpu_arena_step3000.sbatch + +.venv/bin/python experiments/ministral3_3b_followup/merge_eval_json.py \ + --base-json /mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_metrics.json \ + --extra-json /mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_arena_metrics.json \ + --output-json /mnt/vast/runs/andy/deepspec_ministral3_3b_followup/eval_dspark_8gpu_step3000_metrics_complete.json + +.venv/bin/python experiments/ministral3_3b_followup/plot_followup_metrics.py +``` diff --git a/experiments/ministral3_3b_followup/eval_dspark_2gpu_arena_step3000.sbatch b/experiments/ministral3_3b_followup/eval_dspark_2gpu_arena_step3000.sbatch new file mode 100644 index 00000000..e53296e3 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_dspark_2gpu_arena_step3000.sbatch @@ -0,0 +1,48 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds2-arena +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:2 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=160G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/dspark_block7_ministral3_3b_2gpu_arena_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29797} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_dspark_2gpu_step3000_arena_metrics.json" \ + --step 3000 \ + --task arena-hard-v2:500 diff --git a/experiments/ministral3_3b_followup/eval_dspark_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_dspark_3000steps.sbatch new file mode 100644 index 00000000..b05661c5 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_dspark_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=320G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/dspark_block7_ministral3_3b_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29751} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_dspark_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_dspark_4gpu_arena_step3000.sbatch b/experiments/ministral3_3b_followup/eval_dspark_4gpu_arena_step3000.sbatch new file mode 100644 index 00000000..ec728c2f --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_dspark_4gpu_arena_step3000.sbatch @@ -0,0 +1,48 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds4-arena +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:4 +#SBATCH --cpus-per-task=16 +#SBATCH --mem=240G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/dspark_block7_ministral3_3b_4gpu_arena_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29795} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_dspark_4gpu_step3000_arena_metrics.json" \ + --step 3000 \ + --task arena-hard-v2:500 diff --git a/experiments/ministral3_3b_followup/eval_dspark_8gpu_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_dspark_8gpu_3000steps.sbatch new file mode 100644 index 00000000..ea7b2d70 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_dspark_8gpu_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds8-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=320G +#SBATCH --time=02:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/dspark_block7_ministral3_3b_8gpu_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29791} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_dspark_8gpu_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_dspark_8gpu_arena_step3000.sbatch b/experiments/ministral3_3b_followup/eval_dspark_8gpu_arena_step3000.sbatch new file mode 100644 index 00000000..6f1c2aae --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_dspark_8gpu_arena_step3000.sbatch @@ -0,0 +1,48 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds8-arena +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=320G +#SBATCH --time=02:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/dspark_block7_ministral3_3b_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/dspark_block7_ministral3_3b_8gpu_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29793} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_dspark_8gpu_step3000_arena_metrics.json" \ + --step 3000 \ + --task arena-hard-v2:500 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_3000steps.sbatch new file mode 100644 index 00000000..df6ff7d8 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29741} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_8gpu_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_8gpu_3000steps.sbatch new file mode 100644 index 00000000..57df1bc8 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_8gpu_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv8-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_8gpu_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29781} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_8gpu_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_3000steps.sbatch new file mode 100644 index 00000000..110101d9 --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29881} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_from_kl_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_3000steps.sbatch new file mode 100644 index 00000000..aafcffff --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl1e4-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29891} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_from_kl_lr1e4_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch new file mode 100644 index 00000000..ca36211e --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl1e4-8g-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29892} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_from_kl_lr1e4_8gpu_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_ttt5_3000steps.sbatch b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_ttt5_3000steps.sbatch new file mode 100644 index 00000000..949993ec --- /dev/null +++ b/experiments/ministral3_3b_followup/eval_eagle3_e2e_tv_ttt5_3000steps.sbatch @@ -0,0 +1,47 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv5-eval +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=32 +#SBATCH --mem=256G +#SBATCH --time=05:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DRAFT_DIR=/mnt/vast/runs/andy/deepspec_checkpoints/deepspec/eagle3_ttt5_ministral3_3b_e2e_tv/step_latest +TB_DIR=/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/eagle3_ttt5_ministral3_3b_e2e_tv_eval + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29841} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python eval.py \ + --target_name_or_path mistralai/Ministral-3-3B-Instruct-2512 \ + --draft_name_or_path "${DRAFT_DIR}" \ + --tensorboard-dir "${TB_DIR}" \ + --output-json "${RUN_ROOT}/eval_eagle3_e2e_tv_ttt5_step3000_metrics.json" \ + --step 3000 diff --git a/experiments/ministral3_3b_followup/merge_eval_json.py b/experiments/ministral3_3b_followup/merge_eval_json.py new file mode 100644 index 00000000..6a790408 --- /dev/null +++ b/experiments/ministral3_3b_followup/merge_eval_json.py @@ -0,0 +1,79 @@ +import argparse +import json +from pathlib import Path +from typing import Any + + +EXPECTED_DATASETS = [ + "gsm8k", + "math500", + "aime25", + "humaneval", + "mbpp", + "livecodebench", + "mt-bench", + "alpaca", + "arena-hard-v2", +] + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--base-json", type=Path, required=True) + parser.add_argument("--extra-json", type=Path, required=True) + parser.add_argument("--output-json", type=Path, required=True) + return parser.parse_args() + + +def load_json(path: Path) -> dict[str, Any]: + with path.open() as f: + data = json.load(f) + assert isinstance(data, dict), f"Expected JSON object in {path}" + assert "rows" in data, f"Expected rows in {path}" + return data + + +def row_dataset(row: dict[str, Any]) -> str: + dataset = row["dataset"] + assert isinstance(dataset, str), f"Expected string dataset, got {dataset!r}" + return dataset + + +def merge_rows(*row_groups: list[dict[str, Any]]) -> list[dict[str, Any]]: + by_dataset: dict[str, dict[str, Any]] = {} + for rows in row_groups: + for row in rows: + dataset = row_dataset(row) + assert dataset not in by_dataset, f"Duplicate metrics for {dataset}" + by_dataset[dataset] = row + + missing = [dataset for dataset in EXPECTED_DATASETS if dataset not in by_dataset] + assert not missing, f"Missing datasets: {missing}" + return [by_dataset[dataset] for dataset in EXPECTED_DATASETS] + + +def main() -> None: + args = parse_args() + base = load_json(args.base_json) + extra = load_json(args.extra_json) + + for key in ["target_model", "draft_model", "step"]: + assert base[key] == extra[key], ( + f"Mismatched {key}: base={base[key]!r}, extra={extra[key]!r}" + ) + + base_rows = base["rows"] + extra_rows = extra["rows"] + assert isinstance(base_rows, list), f"Expected list rows in {args.base_json}" + assert isinstance(extra_rows, list), f"Expected list rows in {args.extra_json}" + + merged = dict(base) + merged["rows"] = merge_rows(base_rows, extra_rows) + merged["complete"] = True + + args.output_json.parent.mkdir(parents=True, exist_ok=True) + args.output_json.write_text(json.dumps(merged, indent=2) + "\n") + + +if __name__ == "__main__": + main() diff --git a/experiments/ministral3_3b_followup/plot_followup_metrics.py b/experiments/ministral3_3b_followup/plot_followup_metrics.py new file mode 100644 index 00000000..d6727f98 --- /dev/null +++ b/experiments/ministral3_3b_followup/plot_followup_metrics.py @@ -0,0 +1,286 @@ +import argparse +import json +from pathlib import Path + +import matplotlib.pyplot as plt +from tensorboard.backend.event_processing.event_accumulator import EventAccumulator + + +ScalarSeries = list[tuple[int, float]] + +TRAIN_RUNS = { + "kl_baseline": { + "label": "Eagle3 KL baseline", + "tensorboard_dir": Path( + "/mnt/vast/home/andy/tensorboard/deepspec/" + "eagle3_ttt7_ministral3_3b_10k_3000steps" + ), + }, + "e2e_tv_ttt7": { + "label": "Eagle3 e2e-TV TTT-7", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "eagle3_ttt7_ministral3_3b_e2e_tv" + ), + }, + "e2e_tv_ttt5": { + "label": "Eagle3 e2e-TV TTT-5", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "eagle3_ttt5_ministral3_3b_e2e_tv" + ), + }, + "e2e_tv_from_kl": { + "label": "Eagle3 KL -> e2e-TV", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl" + ), + }, + "e2e_tv_from_kl_lr1e4": { + "label": "Eagle3 KL -> e2e-TV lr=1e-4", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4" + ), + }, + "e2e_tv_from_kl_lr1e4_8gpu": { + "label": "Eagle3 KL -> e2e-TV lr=1e-4 8GPU", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "eagle3_ttt7_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu" + ), + }, + "dspark": { + "label": "DSpark block-7", + "tensorboard_dir": Path( + "/mnt/vast/runs/andy/deepspec_tensorboard/deepspec/" + "dspark_block7_ministral3_3b_8gpu" + ), + }, +} + +EVAL_JSONS = { + "kl_baseline": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3/" + "eval_10k_step3000_metrics.json" + ), + "e2e_tv_ttt7": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_eagle3_e2e_tv_step3000_metrics.json" + ), + "e2e_tv_ttt5": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_eagle3_e2e_tv_ttt5_step3000_metrics.json" + ), + "e2e_tv_from_kl": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_eagle3_e2e_tv_from_kl_step3000_metrics.json" + ), + "e2e_tv_from_kl_lr1e4": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_eagle3_e2e_tv_from_kl_lr1e4_step3000_metrics.json" + ), + "e2e_tv_from_kl_lr1e4_8gpu": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_eagle3_e2e_tv_from_kl_lr1e4_8gpu_step3000_metrics.json" + ), + "dspark": Path( + "/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/" + "eval_dspark_8gpu_step3000_metrics.json" + ), +} + + +def load_scalars(tensorboard_dir: Path) -> dict[str, ScalarSeries]: + if not tensorboard_dir.exists(): + return {} + accumulator = EventAccumulator(str(tensorboard_dir), size_guidance={"scalars": 0}) + accumulator.Reload() + + scalars = {} + for tag in accumulator.Tags()["scalars"]: + values_by_step = { + int(event.step): float(event.value) for event in accumulator.Scalars(tag) + } + scalars[tag] = sorted(values_by_step.items()) + return scalars + + +def load_all_scalars() -> dict[str, dict[str, ScalarSeries]]: + return { + run_name: load_scalars(run["tensorboard_dir"]) + for run_name, run in TRAIN_RUNS.items() + } + + +def latest_value(series: ScalarSeries) -> dict[str, float | int] | None: + if not series: + return None + step, value = series[-1] + return {"step": int(step), "value": float(value)} + + +def write_summary( + all_scalars: dict[str, dict[str, ScalarSeries]], + output_path: Path, +) -> None: + summary = {} + for run_name, scalars in sorted(all_scalars.items()): + run_summary = {} + for tag, series in sorted(scalars.items()): + value = latest_value(series) + if value is not None: + run_summary[tag] = value + summary[run_name] = run_summary + + output_path.parent.mkdir(parents=True, exist_ok=True) + output_path.write_text( + json.dumps(summary, indent=2, sort_keys=True, ensure_ascii=False) + "\n" + ) + + +def plot_metric( + *, + all_scalars: dict[str, dict[str, ScalarSeries]], + tag: str, + output_path: Path, + title: str, + ylabel: str, +) -> None: + plt.figure(figsize=(8.5, 4.8)) + for run_name, scalars in all_scalars.items(): + series = scalars.get(tag, []) + if not series: + continue + steps = [step for step, _ in series] + values = [value for _, value in series] + plt.plot( + steps, + values, + label=str(TRAIN_RUNS[run_name]["label"]), + linewidth=1.8, + ) + + plt.title(title) + plt.xlabel("Optimizer step") + plt.ylabel(ylabel) + plt.grid(True, alpha=0.25) + plt.legend(loc="best", fontsize=8) + plt.tight_layout() + output_path.parent.mkdir(parents=True, exist_ok=True) + plt.savefig(output_path, dpi=180) + plt.close() + + +def plot_accept_rates( + *, + all_scalars: dict[str, dict[str, ScalarSeries]], + output_path: Path, +) -> None: + tags = ["train/accept_rate@0", "train/accept_rate@3", "train/accept_rate@6"] + fig, axes = plt.subplots(1, len(tags), figsize=(13, 4.2), sharey=True) + for ax, tag in zip(axes, tags, strict=True): + for run_name, scalars in all_scalars.items(): + series = scalars.get(tag, []) + if not series: + continue + ax.plot( + [step for step, _ in series], + [value for _, value in series], + label=str(TRAIN_RUNS[run_name]["label"]), + linewidth=1.6, + ) + ax.set_title(tag.removeprefix("train/")) + ax.set_xlabel("Optimizer step") + ax.grid(True, alpha=0.25) + axes[0].set_ylabel("Acceptance rate") + axes[-1].legend(loc="best", fontsize=7) + fig.tight_layout() + output_path.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(output_path, dpi=180) + plt.close(fig) + + +def load_eval_acceptance() -> dict[str, dict[str, float]]: + values = {} + for run_name, path in EVAL_JSONS.items(): + if not path.exists(): + continue + payload = json.loads(path.read_text()) + values[run_name] = { + str(row["dataset"]): float(row["acceptance_length"]) + for row in payload["rows"] + } + return values + + +def plot_eval_acceptance(output_path: Path) -> None: + eval_values = load_eval_acceptance() + if len(eval_values) < 2: + return + + datasets = sorted(set().union(*(values.keys() for values in eval_values.values()))) + x_positions = list(range(len(datasets))) + run_names = [run_name for run_name in TRAIN_RUNS if run_name in eval_values] + bar_width = min(0.8 / len(run_names), 0.25) + + plt.figure(figsize=(11, 5.2)) + for idx, run_name in enumerate(run_names): + offset = (idx - (len(run_names) - 1) / 2) * bar_width + values = [ + eval_values[run_name].get(dataset, 0.0) + for dataset in datasets + ] + plt.bar( + [position + offset for position in x_positions], + values, + width=bar_width, + label=str(TRAIN_RUNS[run_name]["label"]), + ) + + plt.xticks(x_positions, datasets, rotation=30, ha="right") + plt.ylabel("Accepted length") + plt.title("Final accepted length by benchmark") + plt.grid(True, axis="y", alpha=0.25) + plt.legend(loc="best", fontsize=8) + plt.tight_layout() + output_path.parent.mkdir(parents=True, exist_ok=True) + plt.savefig(output_path, dpi=180) + plt.close() + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument( + "--output-dir", + type=Path, + default=Path("experiments/ministral3_3b_followup/report_assets"), + ) + args = parser.parse_args() + + all_scalars = load_all_scalars() + write_summary(all_scalars, args.output_dir / "train_metric_summary.json") + plot_metric( + all_scalars=all_scalars, + tag="train/loss", + output_path=args.output_dir / "train_loss.png", + title="Training Loss", + ylabel="Loss", + ) + plot_metric( + all_scalars=all_scalars, + tag="train/tau_probabilistic", + output_path=args.output_dir / "train_tau_probabilistic.png", + title="Train-Side Probabilistic Accepted Length Proxy", + ylabel="Tokens", + ) + plot_accept_rates( + all_scalars=all_scalars, + output_path=args.output_dir / "train_accept_rates.png", + ) + plot_eval_acceptance(args.output_dir / "eval_acceptance_comparison.png") + + +if __name__ == "__main__": + main() diff --git a/experiments/ministral3_3b_followup/report_assets/eval_acceptance_comparison.png b/experiments/ministral3_3b_followup/report_assets/eval_acceptance_comparison.png new file mode 100644 index 00000000..f87b2b48 Binary files /dev/null and b/experiments/ministral3_3b_followup/report_assets/eval_acceptance_comparison.png differ diff --git a/experiments/ministral3_3b_followup/report_assets/train_accept_rates.png b/experiments/ministral3_3b_followup/report_assets/train_accept_rates.png new file mode 100644 index 00000000..47a2f1da Binary files /dev/null and b/experiments/ministral3_3b_followup/report_assets/train_accept_rates.png differ diff --git a/experiments/ministral3_3b_followup/report_assets/train_loss.png b/experiments/ministral3_3b_followup/report_assets/train_loss.png new file mode 100644 index 00000000..8d015a3b Binary files /dev/null and b/experiments/ministral3_3b_followup/report_assets/train_loss.png differ diff --git a/experiments/ministral3_3b_followup/report_assets/train_metric_summary.json b/experiments/ministral3_3b_followup/report_assets/train_metric_summary.json new file mode 100644 index 00000000..f7b273f1 --- /dev/null +++ b/experiments/ministral3_3b_followup/report_assets/train_metric_summary.json @@ -0,0 +1,756 @@ +{ + "dspark": { + "train/accept_rate@0": { + "step": 3000, + "value": 0.9088855385780334 + }, + "train/accept_rate@1": { + "step": 3000, + "value": 0.8897278308868408 + }, + "train/accept_rate@2": { + "step": 3000, + "value": 0.8809123039245605 + }, + "train/accept_rate@3": { + "step": 3000, + "value": 0.8716825246810913 + }, + "train/accept_rate@4": { + "step": 3000, + "value": 0.8612305521965027 + }, + "train/accept_rate@5": { + "step": 3000, + "value": 0.8470444679260254 + }, + "train/accept_rate@6": { + "step": 3000, + "value": 0.8207354545593262 + }, + "train/block_supervised_tokens_abs": { + "step": 3000, + "value": 6.933074951171875 + }, + "train/block_supervised_tokens_ratio": { + "step": 3000, + "value": 0.990439236164093 + }, + "train/ce_loss": { + "step": 3000, + "value": 0.48387807607650757 + }, + "train/confidence_abs_error": { + "step": 3000, + "value": 0.023252103477716446 + }, + "train/confidence_bias": { + "step": 3000, + "value": -5.7178334827767685e-05 + }, + "train/confidence_cumprod_bias": { + "step": 3000, + "value": -6.912671960890293e-05 + }, + "train/confidence_loss": { + "step": 3000, + "value": 0.20054055750370026 + }, + "train/grad_norm": { + "step": 3000, + "value": 0.032470703125 + }, + "train/l1_loss": { + "step": 3000, + "value": 0.23794984817504883 + }, + "train/loss": { + "step": 3000, + "value": 0.421601265668869 + }, + "train/lr": { + "step": 3000, + "value": 0.0 + }, + "train/sampled_anchors_abs": { + "step": 3000, + "value": 395.5830078125 + }, + "train/sampled_anchors_ratio": { + "step": 3000, + "value": 0.7726230621337891 + }, + "train/tau_probabilistic": { + "step": 3000, + "value": 5.878182888031006 + }, + "train/valid_anchors_abs": { + "step": 3000, + "value": 1081.6539306640625 + }, + "train/valid_anchors_ratio": { + "step": 3000, + "value": 0.4878362715244293 + } + }, + "e2e_tv_from_kl": { + "train/accept_rate@0": { + "step": 200, + "value": 0.7781164646148682 + }, + "train/accept_rate@1": { + "step": 200, + "value": 0.7581328749656677 + }, + "train/accept_rate@2": { + "step": 200, + "value": 0.7466332912445068 + }, + "train/accept_rate@3": { + "step": 200, + "value": 0.7343856692314148 + }, + "train/accept_rate@4": { + "step": 200, + "value": 0.7194052934646606 + }, + "train/accept_rate@5": { + "step": 200, + "value": 0.699884295463562 + }, + "train/accept_rate@6": { + "step": 200, + "value": 0.673882246017456 + }, + "train/accuracy@0": { + "step": 200, + "value": 0.8129783272743225 + }, + "train/accuracy@1": { + "step": 200, + "value": 0.7957897782325745 + }, + "train/accuracy@2": { + "step": 200, + "value": 0.7881194949150085 + }, + "train/accuracy@3": { + "step": 200, + "value": 0.7789853811264038 + }, + "train/accuracy@4": { + "step": 200, + "value": 0.7666359543800354 + }, + "train/accuracy@5": { + "step": 200, + "value": 0.7496271133422852 + }, + "train/accuracy@6": { + "step": 200, + "value": 0.7254096269607544 + }, + "train/e2e_tv_loss": { + "step": 200, + "value": 0.4276648163795471 + }, + "train/grad_norm": { + "step": 200, + "value": 0.169921875 + }, + "train/loss": { + "step": 200, + "value": 0.4276648163795471 + }, + "train/lr": { + "step": 200, + "value": 0.0005988585762679577 + }, + "train/tau_greedy": { + "step": 200, + "value": 4.651768207550049 + }, + "train/tau_probabilistic": { + "step": 200, + "value": 4.2134552001953125 + }, + "train/valid_tokens@0": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@1": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@2": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@3": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@4": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@5": { + "step": 200, + "value": 17565.32421875 + }, + "train/valid_tokens@6": { + "step": 200, + "value": 17565.32421875 + } + }, + "e2e_tv_from_kl_lr1e4": { + "train/accept_rate@0": { + "step": 100, + "value": 0.7892638444900513 + }, + "train/accept_rate@1": { + "step": 100, + "value": 0.7840483784675598 + }, + "train/accept_rate@2": { + "step": 100, + "value": 0.782526969909668 + }, + "train/accept_rate@3": { + "step": 100, + "value": 0.7784999012947083 + }, + "train/accept_rate@4": { + "step": 100, + "value": 0.7707067728042603 + }, + "train/accept_rate@5": { + "step": 100, + "value": 0.7571341395378113 + }, + "train/accept_rate@6": { + "step": 100, + "value": 0.7336140275001526 + }, + "train/accuracy@0": { + "step": 100, + "value": 0.8327452540397644 + }, + "train/accuracy@1": { + "step": 100, + "value": 0.8321753144264221 + }, + "train/accuracy@2": { + "step": 100, + "value": 0.8354985117912292 + }, + "train/accuracy@3": { + "step": 100, + "value": 0.8355039954185486 + }, + "train/accuracy@4": { + "step": 100, + "value": 0.8315221071243286 + }, + "train/accuracy@5": { + "step": 100, + "value": 0.8220247030258179 + }, + "train/accuracy@6": { + "step": 100, + "value": 0.8026746511459351 + }, + "train/e2e_tv_loss": { + "step": 100, + "value": 0.39172640442848206 + }, + "train/grad_norm": { + "step": 100, + "value": 0.2734375 + }, + "train/loss": { + "step": 100, + "value": 0.39172640442848206 + }, + "train/lr": { + "step": 100, + "value": 8.416666969424114e-05 + }, + "train/tau_greedy": { + "step": 100, + "value": 5.037959575653076 + }, + "train/tau_probabilistic": { + "step": 100, + "value": 4.44346809387207 + }, + "train/valid_tokens@0": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@1": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@2": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@3": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@4": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@5": { + "step": 100, + "value": 17705.19921875 + }, + "train/valid_tokens@6": { + "step": 100, + "value": 17705.19921875 + } + }, + "e2e_tv_from_kl_lr1e4_8gpu": { + "train/accept_rate@0": { + "step": 960, + "value": 0.8107576966285706 + }, + "train/accept_rate@1": { + "step": 960, + "value": 0.8032916188240051 + }, + "train/accept_rate@2": { + "step": 960, + "value": 0.799018919467926 + }, + "train/accept_rate@3": { + "step": 960, + "value": 0.7926272749900818 + }, + "train/accept_rate@4": { + "step": 960, + "value": 0.782906711101532 + }, + "train/accept_rate@5": { + "step": 960, + "value": 0.7681045532226562 + }, + "train/accept_rate@6": { + "step": 960, + "value": 0.7443888783454895 + }, + "train/accuracy@0": { + "step": 960, + "value": 0.8468327522277832 + }, + "train/accuracy@1": { + "step": 960, + "value": 0.8427723050117493 + }, + "train/accuracy@2": { + "step": 960, + "value": 0.841855525970459 + }, + "train/accuracy@3": { + "step": 960, + "value": 0.8381909728050232 + }, + "train/accuracy@4": { + "step": 960, + "value": 0.8310334086418152 + }, + "train/accuracy@5": { + "step": 960, + "value": 0.8190454840660095 + }, + "train/accuracy@6": { + "step": 960, + "value": 0.798031747341156 + }, + "train/e2e_tv_loss": { + "step": 960, + "value": 0.34104272723197937 + }, + "train/grad_norm": { + "step": 960, + "value": 0.2021484375 + }, + "train/loss": { + "step": 960, + "value": 0.34104272723197937 + }, + "train/lr": { + "step": 960, + "value": 8.043809793889523e-05 + }, + "train/tau_greedy": { + "step": 960, + "value": 5.260736465454102 + }, + "train/tau_probabilistic": { + "step": 960, + "value": 4.743100643157959 + }, + "train/valid_tokens@0": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@1": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@2": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@3": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@4": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@5": { + "step": 960, + "value": 8670.4453125 + }, + "train/valid_tokens@6": { + "step": 960, + "value": 8670.4453125 + } + }, + "e2e_tv_ttt5": { + "train/accept_rate@0": { + "step": 340, + "value": 0.3088866174221039 + }, + "train/accept_rate@1": { + "step": 340, + "value": 0.23770467936992645 + }, + "train/accept_rate@2": { + "step": 340, + "value": 0.19424673914909363 + }, + "train/accept_rate@3": { + "step": 340, + "value": 0.16263574361801147 + }, + "train/accept_rate@4": { + "step": 340, + "value": 0.13550451397895813 + }, + "train/accuracy@0": { + "step": 340, + "value": 0.309736967086792 + }, + "train/accuracy@1": { + "step": 340, + "value": 0.23738116025924683 + }, + "train/accuracy@2": { + "step": 340, + "value": 0.19608260691165924 + }, + "train/accuracy@3": { + "step": 340, + "value": 0.1679745316505432 + }, + "train/accuracy@4": { + "step": 340, + "value": 0.14486446976661682 + }, + "train/e2e_tv_loss": { + "step": 340, + "value": 0.8811959028244019 + }, + "train/grad_norm": { + "step": 340, + "value": 0.0191650390625 + }, + "train/loss": { + "step": 340, + "value": 0.8811959028244019 + }, + "train/lr": { + "step": 340, + "value": 0.0005914028151892126 + }, + "train/tau_greedy": { + "step": 340, + "value": 1.490090012550354 + }, + "train/tau_probabilistic": { + "step": 340, + "value": 1.483672022819519 + }, + "train/valid_tokens@0": { + "step": 340, + "value": 17375.681640625 + }, + "train/valid_tokens@1": { + "step": 340, + "value": 17375.681640625 + }, + "train/valid_tokens@2": { + "step": 340, + "value": 17375.681640625 + }, + "train/valid_tokens@3": { + "step": 340, + "value": 17375.681640625 + }, + "train/valid_tokens@4": { + "step": 340, + "value": 17375.681640625 + } + }, + "e2e_tv_ttt7": { + "train/accept_rate@0": { + "step": 860, + "value": 0.3942939341068268 + }, + "train/accept_rate@1": { + "step": 860, + "value": 0.3225947618484497 + }, + "train/accept_rate@2": { + "step": 860, + "value": 0.27505242824554443 + }, + "train/accept_rate@3": { + "step": 860, + "value": 0.2398243099451065 + }, + "train/accept_rate@4": { + "step": 860, + "value": 0.2107917219400406 + }, + "train/accept_rate@5": { + "step": 860, + "value": 0.18418195843696594 + }, + "train/accept_rate@6": { + "step": 860, + "value": 0.15908914804458618 + }, + "train/accuracy@0": { + "step": 860, + "value": 0.3965262472629547 + }, + "train/accuracy@1": { + "step": 860, + "value": 0.32414039969444275 + }, + "train/accuracy@2": { + "step": 860, + "value": 0.27784740924835205 + }, + "train/accuracy@3": { + "step": 860, + "value": 0.24515438079833984 + }, + "train/accuracy@4": { + "step": 860, + "value": 0.22068458795547485 + }, + "train/accuracy@5": { + "step": 860, + "value": 0.19971370697021484 + }, + "train/accuracy@6": { + "step": 860, + "value": 0.18134939670562744 + }, + "train/e2e_tv_loss": { + "step": 860, + "value": 0.8688939809799194 + }, + "train/grad_norm": { + "step": 860, + "value": 0.0220947265625 + }, + "train/loss": { + "step": 860, + "value": 0.8688939809799194 + }, + "train/lr": { + "step": 860, + "value": 0.0005074540385976434 + }, + "train/tau_greedy": { + "step": 860, + "value": 1.7387313842773438 + }, + "train/tau_probabilistic": { + "step": 860, + "value": 1.7238452434539795 + }, + "train/valid_tokens@0": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@1": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@2": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@3": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@4": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@5": { + "step": 860, + "value": 17503.728515625 + }, + "train/valid_tokens@6": { + "step": 860, + "value": 17503.728515625 + } + }, + "kl_baseline": { + "train/accept_rate@0": { + "step": 3000, + "value": 0.7828433513641357 + }, + "train/accept_rate@1": { + "step": 3000, + "value": 0.7800154685974121 + }, + "train/accept_rate@2": { + "step": 3000, + "value": 0.7808699011802673 + }, + "train/accept_rate@3": { + "step": 3000, + "value": 0.7791018486022949 + }, + "train/accept_rate@4": { + "step": 3000, + "value": 0.7734428644180298 + }, + "train/accept_rate@5": { + "step": 3000, + "value": 0.7620912194252014 + }, + "train/accept_rate@6": { + "step": 3000, + "value": 0.740893542766571 + }, + "train/accuracy@0": { + "step": 3000, + "value": 0.8395382761955261 + }, + "train/accuracy@1": { + "step": 3000, + "value": 0.8423240184783936 + }, + "train/accuracy@2": { + "step": 3000, + "value": 0.848247230052948 + }, + "train/accuracy@3": { + "step": 3000, + "value": 0.851127028465271 + }, + "train/accuracy@4": { + "step": 3000, + "value": 0.8497721552848816 + }, + "train/accuracy@5": { + "step": 3000, + "value": 0.8434146046638489 + }, + "train/accuracy@6": { + "step": 3000, + "value": 0.8280242085456848 + }, + "train/grad_norm": { + "step": 3000, + "value": 0.078125 + }, + "train/loss": { + "step": 3000, + "value": 2.1320221424102783 + }, + "train/lr": { + "step": 3000, + "value": 0.0 + }, + "train/ploss_0": { + "step": 3000, + "value": 0.5350198745727539 + }, + "train/ploss_1": { + "step": 3000, + "value": 0.535610020160675 + }, + "train/ploss_2": { + "step": 3000, + "value": 0.5331065654754639 + }, + "train/ploss_3": { + "step": 3000, + "value": 0.5346406698226929 + }, + "train/ploss_4": { + "step": 3000, + "value": 0.540634036064148 + }, + "train/ploss_5": { + "step": 3000, + "value": 0.552795946598053 + }, + "train/ploss_6": { + "step": 3000, + "value": 0.5760428309440613 + }, + "train/tau_greedy": { + "step": 3000, + "value": 5.192190647125244 + }, + "train/tau_probabilistic": { + "step": 3000, + "value": 4.4316630363464355 + }, + "train/valid_tokens@0": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@1": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@2": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@3": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@4": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@5": { + "step": 3000, + "value": 8655.0615234375 + }, + "train/valid_tokens@6": { + "step": 3000, + "value": 8655.0615234375 + } + } +} diff --git a/experiments/ministral3_3b_followup/report_assets/train_tau_probabilistic.png b/experiments/ministral3_3b_followup/report_assets/train_tau_probabilistic.png new file mode 100644 index 00000000..ccfed0e6 Binary files /dev/null and b/experiments/ministral3_3b_followup/report_assets/train_tau_probabilistic.png differ diff --git a/experiments/ministral3_3b_followup/train_dspark_3000steps_16gpu.sbatch b/experiments/ministral3_3b_followup/train_dspark_3000steps_16gpu.sbatch new file mode 100644 index 00000000..a397aeda --- /dev/null +++ b/experiments/ministral3_3b_followup/train_dspark_3000steps_16gpu.sbatch @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds-3k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --nodes=2 +#SBATCH --ntasks=2 +#SBATCH --ntasks-per-node=1 +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=04:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +export REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +export RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +export DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +export CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR +MASTER_ADDR=$(scontrol show hostnames "${SLURM_JOB_NODELIST}" | head -n 1) +export MASTER_PORT=${MASTER_PORT:-29731} +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +srun --ntasks="${SLURM_NNODES}" --ntasks-per-node=1 bash -lc ' +set -euo pipefail +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export RANK="${SLURM_PROCID}" +export WORLD_SIZE="${SLURM_NNODES}" + +python train.py \ + --config config/dspark/dspark_ministral3_3b.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" +' diff --git a/experiments/ministral3_3b_followup/train_dspark_3000steps_8gpu.sbatch b/experiments/ministral3_3b_followup/train_dspark_3000steps_8gpu.sbatch new file mode 100644 index 00000000..c2e5a8c9 --- /dev/null +++ b/experiments/ministral3_3b_followup/train_dspark_3000steps_8gpu.sbatch @@ -0,0 +1,52 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds-8g +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=00:10:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29771} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/dspark/dspark_ministral3_3b.py \ + --opts "exp_name=dspark_block7_ministral3_3b_8gpu" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=10" \ + --opts "logging.keep_last_checkpoints=8" diff --git a/experiments/ministral3_3b_followup/train_dspark_smoke.sbatch b/experiments/ministral3_3b_followup/train_dspark_smoke.sbatch new file mode 100644 index 00000000..2c4b1fcf --- /dev/null +++ b/experiments/ministral3_3b_followup/train_dspark_smoke.sbatch @@ -0,0 +1,51 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-ds-smoke +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:1 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=120G +#SBATCH --time=00:40:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${DATA_ROOT}/smoke_target_cache + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29711} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/dspark/dspark_ministral3_3b.py \ + --opts "exp_name=dspark_block7_ministral3_3b_smoke" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=0" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=1" \ + --opts "train.max_train_steps=2" \ + --opts "logging.checkpointing_steps=1" \ + --opts "logging.keep_last_checkpoints=2" diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_16gpu.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_16gpu.sbatch new file mode 100644 index 00000000..8213f6df --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_16gpu.sbatch @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv-3k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --nodes=2 +#SBATCH --ntasks=2 +#SBATCH --ntasks-per-node=1 +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +export REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +export RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +export DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +export CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR +MASTER_ADDR=$(scontrol show hostnames "${SLURM_JOB_NODELIST}" | head -n 1) +export MASTER_PORT=${MASTER_PORT:-29721} +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +srun --ntasks="${SLURM_NNODES}" --ntasks-per-node=1 bash -lc ' +set -euo pipefail +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export RANK="${SLURM_PROCID}" +export WORLD_SIZE="${SLURM_NNODES}" + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" +' diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_8gpu.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_8gpu.sbatch new file mode 100644 index 00000000..d00fc29f --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_3000steps_8gpu.sbatch @@ -0,0 +1,52 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv-8g +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29761} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv.py \ + --opts "exp_name=eagle3_ttt7_ministral3_3b_e2e_tv_8gpu" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_3000steps_16gpu.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_3000steps_16gpu.sbatch new file mode 100644 index 00000000..d54a1ee6 --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_3000steps_16gpu.sbatch @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl-3k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --nodes=2 +#SBATCH --ntasks=2 +#SBATCH --ntasks-per-node=1 +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +export REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +export RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +export DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +export CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR +MASTER_ADDR=$(scontrol show hostnames "${SLURM_JOB_NODELIST}" | head -n 1) +export MASTER_PORT=${MASTER_PORT:-29861} +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +srun --ntasks="${SLURM_NNODES}" --ntasks-per-node=1 bash -lc ' +set -euo pipefail +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export RANK="${SLURM_PROCID}" +export WORLD_SIZE="${SLURM_NNODES}" + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" +' diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_3000steps_16gpu.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_3000steps_16gpu.sbatch new file mode 100644 index 00000000..bb0d84e3 --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_3000steps_16gpu.sbatch @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl1e4 +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --nodes=2 +#SBATCH --ntasks=2 +#SBATCH --ntasks-per-node=1 +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +export REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +export RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +export DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +export CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR +MASTER_ADDR=$(scontrol show hostnames "${SLURM_JOB_NODELIST}" | head -n 1) +export MASTER_PORT=${MASTER_PORT:-29863} +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +srun --ntasks="${SLURM_NNODES}" --ntasks-per-node=1 bash -lc ' +set -euo pipefail +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export RANK="${SLURM_PROCID}" +export WORLD_SIZE="${SLURM_NNODES}" + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" +' diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch new file mode 100644 index 00000000..25323e3d --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_from_kl_lr1e4_8gpu_3000steps.sbatch @@ -0,0 +1,51 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tvkl1e4-8g +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=00:30:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29864} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv_from_kl_lr1e4_8gpu.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_smoke.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_smoke.sbatch new file mode 100644 index 00000000..50eb1b85 --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_smoke.sbatch @@ -0,0 +1,51 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv-smoke +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --gres=gpu:h100:1 +#SBATCH --cpus-per-task=8 +#SBATCH --mem=100G +#SBATCH --time=00:30:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +CACHE_DIR=${DATA_ROOT}/smoke_target_cache + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR=127.0.0.1 +export MASTER_PORT=${MASTER_PORT:-29701} +export RANK=0 +export WORLD_SIZE=1 +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv.py \ + --opts "exp_name=eagle3_ttt7_ministral3_3b_e2e_tv_smoke" \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=0" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=1" \ + --opts "train.max_train_steps=2" \ + --opts "logging.checkpointing_steps=1" \ + --opts "logging.keep_last_checkpoints=2" diff --git a/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_ttt5_3000steps_16gpu.sbatch b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_ttt5_3000steps_16gpu.sbatch new file mode 100644 index 00000000..b7dc2b5e --- /dev/null +++ b/experiments/ministral3_3b_followup/train_eagle3_e2e_tv_ttt5_3000steps_16gpu.sbatch @@ -0,0 +1,58 @@ +#!/usr/bin/env bash +#SBATCH --job-name=dspec-min3-tv5-3k +#SBATCH --partition=h100 +#SBATCH --qos=research +#SBATCH --nodes=2 +#SBATCH --ntasks=2 +#SBATCH --ntasks-per-node=1 +#SBATCH --gres=gpu:h100:8 +#SBATCH --cpus-per-task=64 +#SBATCH --mem=512G +#SBATCH --time=03:00:00 +#SBATCH --output=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.out +#SBATCH --error=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup/logs/%x-%j.err + +set -euo pipefail + +export REPO_DIR=/mnt/vast/home/andy/code/DeepSpec-ministral-repro-20260702-151809 +export RUN_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_followup +export DATA_ROOT=/mnt/vast/runs/andy/deepspec_ministral3_3b_eagle3 +export CACHE_DIR=${DATA_ROOT}/target_cache_10k + +mkdir -p "${RUN_ROOT}/logs" +mkdir -p /mnt/vast/runs/andy/hf_home +mkdir -p /mnt/vast/runs/andy/torch_home +mkdir -p /mnt/vast/runs/andy/triton_cache +mkdir -p /mnt/vast/runs/andy/xdg_cache + +export HF_HOME=/mnt/vast/runs/andy/hf_home +export HF_DATASETS_CACHE=${HF_HOME}/datasets +export TRANSFORMERS_CACHE=${HF_HOME}/transformers +export TORCH_HOME=/mnt/vast/runs/andy/torch_home +export TRITON_CACHE_DIR=/mnt/vast/runs/andy/triton_cache/${SLURM_JOB_ID} +export XDG_CACHE_HOME=/mnt/vast/runs/andy/xdg_cache +export MASTER_ADDR +MASTER_ADDR=$(scontrol show hostnames "${SLURM_JOB_NODELIST}" | head -n 1) +export MASTER_PORT=${MASTER_PORT:-29821} +export TOKENIZERS_PARALLELISM=false +export USE_HUB_KERNELS=0 + +srun --ntasks="${SLURM_NNODES}" --ntasks-per-node=1 bash -lc ' +set -euo pipefail +cd "${REPO_DIR}" +source .venv/bin/activate +export PYTHONPATH="${REPO_DIR}:${PYTHONPATH:-}" +export RANK="${SLURM_PROCID}" +export WORLD_SIZE="${SLURM_NNODES}" + +python train.py \ + --config config/eagle3/eagle3_ministral3_3b_e2e_tv_ttt5.py \ + --opts "data.target_cache_path=${CACHE_DIR}" \ + --opts "data.num_workers=2" \ + --opts "train.local_batch_size=1" \ + --opts "train.global_batch_size=512" \ + --opts "train.max_train_steps=3000" \ + --opts "logging.checkpointing_steps=3000" \ + --opts "logging.save_only_checkpointing_steps=25" \ + --opts "logging.keep_last_checkpoints=8" +' diff --git a/requirements.txt b/requirements.txt index 30316637..c8515bfe 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,6 +2,8 @@ # is not appropriate for your environment. torch==2.9.1 transformers==5.10.2 +accelerate==1.14.0 +kernels==0.14.1 numpy==2.4.4 PyYAML==6.0.3 tqdm==4.67.3 diff --git a/scripts/data/download_streaming_subset.py b/scripts/data/download_streaming_subset.py new file mode 100644 index 00000000..ad0d88b6 --- /dev/null +++ b/scripts/data/download_streaming_subset.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +import argparse +import json +import os +import sys +from pathlib import Path + +from datasets import load_dataset + + +ROLE_MAPPING = { + "human": "user", + "gpt": "assistant", + "chatgpt": "assistant", + "bing": "assistant", + "bard": "assistant", +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Write the first N Open-PerfectBlend rows as DeepSpec JSONL." + ) + parser.add_argument("--dataset-name", default="mlabonne/open-perfectblend") + parser.add_argument("--split", default="train") + parser.add_argument("--sample-size", type=int, required=True) + parser.add_argument("--output-path", type=Path, required=True) + parser.add_argument("--skip-existing", action="store_true") + return parser.parse_args() + + +def normalize_conversations(row: dict, idx: int) -> dict: + conversations = [] + for message in row["conversations"]: + role = ROLE_MAPPING.get(message["from"]) + if role is None: + continue + conversations.append({"role": role, "content": message["value"]}) + assert conversations, f"row {idx} produced no conversations." + assert conversations[0]["role"] == "user", ( + f"row {idx} does not start with user: {conversations[0]['role']}" + ) + return {"id": idx, "conversations": conversations} + + +def main() -> None: + args = parse_args() + assert args.sample_size > 0, f"sample_size must be positive, got {args.sample_size}" + if args.output_path.exists(): + if args.skip_existing: + print(f"skip existing output: {args.output_path}") + return + raise FileExistsError(f"Output JSONL already exists: {args.output_path}") + + args.output_path.parent.mkdir(parents=True, exist_ok=True) + dataset = load_dataset(args.dataset_name, split=args.split, streaming=True) + written = 0 + with args.output_path.open("w", encoding="utf-8") as handle: + for idx, row in enumerate(dataset): + converted = normalize_conversations(row, idx) + handle.write(json.dumps(converted, ensure_ascii=False) + "\n") + written += 1 + if written >= args.sample_size: + break + print(f"wrote {written} rows -> {args.output_path}", flush=True) + sys.stdout.flush() + os._exit(0) + + +if __name__ == "__main__": + main() diff --git a/scripts/data/generate_train_data_transformers.py b/scripts/data/generate_train_data_transformers.py new file mode 100644 index 00000000..39c77071 --- /dev/null +++ b/scripts/data/generate_train_data_transformers.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path +from time import perf_counter + +os.environ["USE_HUB_KERNELS"] = "0" + +import torch +from tqdm import tqdm + +from deepspec.data.parser import encode_chat_messages +from deepspec.modeling.target_utils import load_target_causal_lm, load_target_tokenizer + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Regenerate DeepSpec training conversations with a local HF model." + ) + parser.add_argument("--model", required=True) + parser.add_argument("--input-file-path", type=Path, required=True) + parser.add_argument("--output-file-path", type=Path, required=True) + parser.add_argument("--error-file-path", type=Path) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument("--max-samples", type=int) + parser.add_argument("--resume", action="store_true") + parser.add_argument("--temperature", type=float, default=0.7) + parser.add_argument("--top-p", type=float, default=0.8) + parser.add_argument("--top-k", type=int, default=20) + parser.add_argument("--max-new-tokens", type=int, default=4096) + parser.add_argument("--attn-implementation", default="sdpa") + return parser.parse_args() + + +def count_lines(path: Path) -> int: + if not path.exists(): + return 0 + with path.open("r", encoding="utf-8") as handle: + return sum(1 for _ in handle) + + +def iter_shard_rows(path: Path, shard_index: int, num_shards: int): + with path.open("r", encoding="utf-8") as handle: + for line_index, line in enumerate(handle): + if line_index % num_shards != shard_index: + continue + yield line_index, json.loads(line) + + +def generate_assistant( + *, + model, + tokenizer, + messages: list[dict], + args: argparse.Namespace, +) -> str: + input_ids = encode_chat_messages( + tokenizer, + messages, + add_generation_prompt=True, + ).to(model.device) + attention_mask = torch.ones_like(input_ids) + do_sample = float(args.temperature) > 0.0 + generation_kwargs = { + "input_ids": input_ids, + "attention_mask": attention_mask, + "do_sample": do_sample, + "max_new_tokens": int(args.max_new_tokens), + "eos_token_id": tokenizer.eos_token_id, + "pad_token_id": tokenizer.pad_token_id, + } + if do_sample: + generation_kwargs.update( + { + "temperature": float(args.temperature), + "top_p": float(args.top_p), + "top_k": int(args.top_k), + } + ) + with torch.inference_mode(): + output_ids = model.generate(**generation_kwargs) + generated_ids = output_ids[0, input_ids.shape[1] :] + return tokenizer.decode(generated_ids, skip_special_tokens=True).strip() + + +def regenerate_sample(*, model, tokenizer, sample: dict, args: argparse.Namespace) -> dict: + regenerated = [] + for message in sample["conversations"]: + role = message["role"] + if role == "system": + regenerated.append(message) + continue + if role == "assistant": + continue + assert role == "user", f"Unsupported message role: {role}" + regenerated.append(message) + regenerated.append( + { + "role": "assistant", + "content": generate_assistant( + model=model, + tokenizer=tokenizer, + messages=regenerated, + args=args, + ), + } + ) + sample = dict(sample) + sample["conversations"] = regenerated + sample["status"] = "success" + return sample + + +def main() -> None: + args = parse_args() + assert 0 <= args.shard_index < args.num_shards, ( + f"Expected shard_index in [0, {args.num_shards}), got {args.shard_index}" + ) + assert args.max_new_tokens > 0, ( + f"max_new_tokens must be positive, got {args.max_new_tokens}" + ) + args.output_file_path.parent.mkdir(parents=True, exist_ok=True) + error_path = args.error_file_path + if error_path is None: + error_path = args.output_file_path.with_name( + args.output_file_path.stem + "_error.jsonl" + ) + error_path.parent.mkdir(parents=True, exist_ok=True) + + completed = count_lines(args.output_file_path) + count_lines(error_path) + if not args.resume: + completed = 0 + + tokenizer = load_target_tokenizer(args.model) + model = load_target_causal_lm( + args.model, + dtype=torch.bfloat16, + attn_implementation=args.attn_implementation, + ).to(device="cuda").eval() + + total_seen = 0 + success = 0 + errors = 0 + start_time = perf_counter() + output_mode = "a" if args.resume else "w" + with ( + args.output_file_path.open(output_mode, encoding="utf-8") as output_handle, + error_path.open(output_mode, encoding="utf-8") as error_handle, + ): + rows = iter_shard_rows(args.input_file_path, args.shard_index, args.num_shards) + progress = tqdm(rows, desc=f"shard {args.shard_index}/{args.num_shards}") + for _line_index, sample in progress: + if total_seen < completed: + total_seen += 1 + continue + if args.max_samples is not None and success + errors >= args.max_samples: + break + try: + regenerated = regenerate_sample( + model=model, + tokenizer=tokenizer, + sample=sample, + args=args, + ) + except Exception as exc: + sample = dict(sample) + sample["status"] = "error" + sample["error"] = str(exc) + error_handle.write(json.dumps(sample, ensure_ascii=False) + "\n") + error_handle.flush() + errors += 1 + else: + output_handle.write(json.dumps(regenerated, ensure_ascii=False) + "\n") + output_handle.flush() + success += 1 + total_seen += 1 + elapsed = max(perf_counter() - start_time, 1e-6) + progress.set_postfix(success=success, errors=errors, samples_per_s=success / elapsed) + + print(f"success={success} errors={errors} output={args.output_file_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts/data/prepare_target_cache.py b/scripts/data/prepare_target_cache.py index d072582b..9e76a569 100644 --- a/scripts/data/prepare_target_cache.py +++ b/scripts/data/prepare_target_cache.py @@ -3,10 +3,12 @@ import json import os +os.environ["USE_HUB_KERNELS"] = "0" + import torch import torch.distributed as dist from torch.utils.data import DataLoader, Subset -from transformers import AutoModel, AutoTokenizer +from transformers import AutoModel from deepspec.data import ConversationCollator from deepspec.data.target_cache_dataset import ( @@ -23,6 +25,11 @@ rename_local_target_cache_shards, write_target_cache_manifest, ) +from deepspec.modeling.target_utils import ( + get_target_backbone, + get_target_hidden_size, + load_target_tokenizer, +) from deepspec.data.jsonl_dataset import JsonLineDataset from deepspec.utils import ( CustomJSONEncoder, @@ -52,24 +59,6 @@ class TargetForwardResult: target_last_hidden_states: torch.Tensor -def _get_target_backbone(target_model): - model_type = str(target_model.config.model_type) - if model_type in ("gemma4", "gemma4_unified"): - if hasattr(target_model, "language_model"): - return target_model.language_model - if hasattr(target_model, "model") and hasattr(target_model.model, "language_model"): - return target_model.model.language_model - assert False, "Gemma4 target model must expose a text language_model." - return getattr(target_model, "model", target_model) - - -def _get_target_hidden_size(target_model) -> int: - model_type = str(target_model.config.model_type) - if model_type in ("gemma4", "gemma4_unified"): - return int(target_model.config.text_config.hidden_size) - return int(target_model.config.hidden_size) - - def _get_hook_tensor(output): if isinstance(output, torch.Tensor): return output @@ -87,7 +76,7 @@ def run_target_forward_with_hooks( attention_mask: torch.Tensor, target_layer_ids, ): - backbone = _get_target_backbone(target_model) + backbone = get_target_backbone(target_model) layer_modules = backbone.layers target_layer_ids = [int(layer_id) for layer_id in target_layer_ids] captured_hidden_states = {} @@ -248,7 +237,7 @@ def main(local_rank: int): local_total_samples = local_end - local_start local_subset = Subset(dataset, range(local_start, local_end)) - tokenizer = AutoTokenizer.from_pretrained( + tokenizer = load_target_tokenizer( config.model.target_model_name_or_path, ) target_model = AutoModel.from_pretrained( @@ -256,7 +245,7 @@ def main(local_rank: int): dtype=torch.bfloat16, attn_implementation="sdpa", ).to(device=device).eval() - target_hidden_size = _get_target_hidden_size(target_model) + target_hidden_size = get_target_hidden_size(target_model) train_collator = ConversationCollator( tokenizer=tokenizer, chat_template=config.data.chat_template, diff --git a/train.py b/train.py index 40d9e074..bcb4776a 100644 --- a/train.py +++ b/train.py @@ -1,6 +1,9 @@ import argparse import json import os + +os.environ["USE_HUB_KERNELS"] = "0" + import torch from deepspec.utils import ( CustomJSONEncoder,