Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions examples/post_training/modelopt/convert_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,12 +18,13 @@
from megatron.core.parallel_state import destroy_model_parallel
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import report_current_memory_info, to_empty_if_meta
from megatron.training import get_args, get_tokenizer
from megatron.training.checkpointing import save_checkpoint
from megatron.training.initialize import initialize_megatron
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

ALGO_TO_CONFIG = {
"eagle1": mtsp.config.EAGLE1_DEFAULT_CFG,
Expand Down Expand Up @@ -120,13 +121,13 @@ def check_arguments():
UserWarning,
)

model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)
report_current_memory_info()

unwrapped_model = unwrap_model(model)[0]

if args.pretrained_model_path is not None:
import_dtype = torch.float16 if args.fp16 else torch.bfloat16
import_dtype = torch.float16 if args.fp16 else torch.bfloat16
unwrapped_model = unwrap_model(model)[0]
workspace_dir = os.environ.get("MLM_WORK_DIR", "/tmp")
print_rank_0("Import model from Hugging Face checkpoint in dtype {}.".format(str(import_dtype)))
Expand Down Expand Up @@ -168,10 +169,10 @@ def check_arguments():
tokenizer = get_tokenizer()
for i in range(unwrapped_model.eagle_config.parallel_draft_step - 1):
mask_token = "[MASK_{}]".format(i)
tokenizer._tokenizer.add_tokens([mask_token], special_tokens=True)
tokenizer._tokenizer.add_tokens([mask_token], special_tokens=True)
token_id = tokenizer._tokenizer.convert_tokens_to_ids(mask_token)
setattr(unwrapped_model, "mask_token_{}".format(i), torch.tensor(token_id))

elif args.algorithm == "medusa":
config = {"medusa_num_heads": args.export_num_medusa_heads, "medusa_num_layers": 1}
unwrapped_model = mtsp.convert(unwrapped_model, [("medusa", config)])
Expand Down
9 changes: 5 additions & 4 deletions examples/post_training/modelopt/export.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,10 +13,11 @@

from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.training import get_args, get_model
from megatron.training.initialize import initialize_megatron
from megatron.training.utils import unwrap_model
from model_provider import model_provider

warnings.filterwarnings('ignore')

Expand Down Expand Up @@ -64,11 +65,11 @@ def add_modelopt_export_args(parser):
UserWarning,
)

model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)

# Materialize the model from meta device to cpu before loading the checkpoint.
# Materialize the model from meta device to cpu before loading the checkpoint.
unwrapped_model = unwrap_model(model)[0]
unwrapped_model.to_empty(device="cpu")
unwrapped_model.to_empty(device="cpu")

if args.load is not None:
_ = load_modelopt_checkpoint(model)
Expand Down
15 changes: 7 additions & 8 deletions examples/post_training/modelopt/finetune.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,12 @@

"""Supervised Finetuning GPT."""
import itertools
import json
import os
import sys
from functools import partial
from typing import Any, Dict, Optional

import json
import jsonlines

sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../")))
Expand All @@ -20,17 +20,16 @@
from megatron.core.enums import ModelType
from megatron.core.models.gpt import GPTModel
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.model_provider import model_provider
from megatron.post_training.loss_func import loss_func
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.non_loss_data_func import report_draft_acceptance_length
from megatron.training import get_args, get_timers, get_tokenizer, pretrain
from megatron.training.utils import (
average_losses_across_data_parallel_group,
get_batch_on_this_cp_rank,
get_ltor_masks_and_position_ids,
print_rank_0,
unwrap_model,
)
from model_provider import model_provider

REMOVE_THINK_CHAT_TEMPLATE = (
"{% if '</think>' in content %}{% set content = content.split('</think>')[-1] %}{% endif %}"
Expand All @@ -48,7 +47,7 @@ def add_finetune_args(parser):

def get_eos_id():
"""Return the eos token id.

We insert eos_token between two samples during packing. However, if the eos_token is used in message or after turns,
we need to replace it with some other special tokens that do not appear in message."""
tokenizer = get_tokenizer()
Expand Down Expand Up @@ -79,7 +78,7 @@ def __init__(self, data_dir: str, num_samples):

def __len__(self):
return self.num_samples

def __getitem__(self, idx):
idx = idx % len(self.file_paths)
file_path = self.file_paths[idx]
Expand Down Expand Up @@ -387,7 +386,7 @@ def train_valid_test_sft_datasets_provider(train_val_test_num_samples):

def get_batch(data_iterator):
"""Generate a batch.

For OfflineDataset, the aux_hidden_states and final hidden_states from the
base model are loaded for offline speculative model training."""
# TODO: this is pretty hacky, find a better way
Expand Down Expand Up @@ -495,7 +494,7 @@ def forward_step(data_iterator, model: GPTModel):
if __name__ == "__main__":
pretrain(
train_valid_test_sft_datasets_provider,
model_provider,
partial(model_provider, modelopt_gpt_mamba_builder),
ModelType.encoder_or_decoder,
forward_step,
extra_args_provider=add_finetune_args,
Expand Down
5 changes: 3 additions & 2 deletions examples/post_training/modelopt/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,11 @@
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.generate import simple_generate
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import report_current_memory_info, to_empty_if_meta
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

warnings.filterwarnings('once')

Expand Down Expand Up @@ -96,7 +97,7 @@ def get_conversations(example):
UserWarning,
)

model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)
report_current_memory_info()

unwrapped_model = unwrap_model(model)[0]
Expand Down
26 changes: 10 additions & 16 deletions examples/post_training/modelopt/generation_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,18 @@
import os
import sys
import warnings
from functools import partial

sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../")))
import os
import sys
from argparse import Namespace
from contextlib import nullcontext

import torch

from megatron.core import mpu
from megatron.core.inference.engines import AbstractEngine, StaticInferenceEngine
from megatron.core.inference.engines.abstract_engine import AbstractEngine
from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import (
InferenceWrapperConfig,
Expand All @@ -19,31 +24,19 @@
from megatron.core.inference.text_generation_controllers.text_generation_controller import (
TextGenerationController,
)
import torch

from megatron.core.inference.engines import AbstractEngine, StaticInferenceEngine
from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import (
InferenceWrapperConfig,
)
from megatron.training import get_model
from megatron.core.transformer.module import MegatronModule
from megatron.inference.text_generation import beam_search_and_post_process
from megatron.inference.text_generation.mcore_engine_server import (
ModelInferenceWrapperServer,
run_mcore_engine,
)
from megatron.inference.text_generation_server import MegatronServer
from megatron.training import print_rank_0

from megatron.core import mpu
from megatron.training import get_args, get_model, get_tokenizer
from megatron.post_training.arguments import add_modelopt_args
from megatron.training import get_args, get_model, get_tokenizer, print_rank_0
from megatron.training.checkpointing import load_checkpoint
from megatron.training.initialize import initialize_megatron


from megatron.post_training.arguments import add_modelopt_args


def get_inference_engine(args: Namespace, model: MegatronModule) -> AbstractEngine:
"""Get the relevant backend for running inference

Expand Down Expand Up @@ -144,9 +137,10 @@ def main(model_provider: str = "gpt"):
load_context = fp8_model_init()
with load_context:

from megatron.post_training.model_provider import model_provider as modelopt_model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from model_provider import model_provider as root_model_provider
if model_provider == "gpt":
model = get_model(modelopt_model_provider, wrap_with_ddp=False)
model = get_model(partial(root_model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)
elif model_provider == "mamba":
pass
else:
Expand Down
5 changes: 3 additions & 2 deletions examples/post_training/modelopt/mmlu.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,11 @@
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.generate import simple_generate
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import report_current_memory_info
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

warnings.filterwarnings('ignore')

Expand Down Expand Up @@ -148,7 +149,7 @@ def generate_prompt(test_example, dev_examples, few_shots=0, no_subject_prompt=F
UserWarning,
)

model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)
report_current_memory_info()

disable_tqdm = args.disable_tqdm or torch.distributed.get_rank() > 0
Expand Down
10 changes: 6 additions & 4 deletions examples/post_training/modelopt/offline_feature_extract.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,19 +4,21 @@
import functools
import os
import sys

import torch

sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../")))


from examples.post_training.modelopt.finetune import SFTDataset
from megatron.core import mpu
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

from examples.post_training.modelopt.finetune import SFTDataset

def add_extract_args(parser):
"""Add additional arguments for feature extraction."""
Expand Down Expand Up @@ -51,7 +53,7 @@ def extract_feature(dataset, model, output_dir, idx_start, idx_end):

args = get_args()
tokenizer = get_tokenizer()
model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)

load_modelopt_checkpoint(model, strict=not args.untie_embeddings_and_output_weights)
print_rank_0("Done loading checkpoint")
Expand All @@ -68,7 +70,7 @@ def extract_feature(dataset, model, output_dir, idx_start, idx_end):
"shard_index": mpu.get_expert_data_parallel_rank(),
}
sft_dataset = SFTDataset(args.num_samples, None, **kwargs)

extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "train"), 0, int(args.num_samples * 0.98))
extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "valid"), int(args.num_samples * 0.98), int(args.num_samples * 0.99))
extract_feature(sft_dataset, unwrapped_model, os.path.join(args.output_dir, "test"), int(args.num_samples * 0.99), args.num_samples)
Expand Down
5 changes: 3 additions & 2 deletions examples/post_training/modelopt/prune.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,11 +23,12 @@
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.generate import simple_generate
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import report_current_memory_info
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.checkpointing import save_checkpoint
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

warnings.filterwarnings("ignore")

Expand Down Expand Up @@ -138,7 +139,7 @@ def get_calib_dataloader(calib_size=1024, max_sequence_length=512):
check_arguments(args)

tokenizer = get_tokenizer()._tokenizer
model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)
unwrapped_model = unwrap_model(model)[0]

report_current_memory_info()
Expand Down
5 changes: 3 additions & 2 deletions examples/post_training/modelopt/quantize.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,11 +20,12 @@
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.generate import simple_generate
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import report_current_memory_info
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.checkpointing import save_checkpoint
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

warnings.filterwarnings("ignore")

Expand Down Expand Up @@ -157,7 +158,7 @@ def get_calib_dataloader(calib_size=512, max_sequence_length=512):
args = get_args()

tokenizer = get_tokenizer()._tokenizer
model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)

report_current_memory_info()

Expand Down
18 changes: 6 additions & 12 deletions examples/post_training/modelopt/validate.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,29 +2,23 @@

"""Sample Generate GPT."""
import functools
import json
import os
import sys
import warnings
import json

sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../")))

import modelopt
from modelopt.torch.speculative.plugins.megatron_eagle import MegatronARValidation
import torch
from datasets import load_dataset
from tqdm import tqdm
from modelopt.torch.speculative.plugins.megatron_eagle import MegatronARValidation

from megatron.core import mpu
from megatron.core.inference.communication_utils import broadcast_from_last_pipeline_stage
from megatron.core.pipeline_parallel import get_forward_backward_func
from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
from megatron.post_training.model_provider import model_provider
from megatron.post_training.model_builder import modelopt_gpt_mamba_builder
from megatron.post_training.utils import get_mtbench_chat_data
from megatron.training import get_args, get_model, get_tokenizer, initialize_megatron
from megatron.training.utils import get_ltor_masks_and_position_ids, print_rank_0, unwrap_model
from megatron.training.utils import print_rank_0, unwrap_model
from model_provider import model_provider

warnings.filterwarnings('ignore')

Expand Down Expand Up @@ -121,7 +115,7 @@ def report_current_memory_info():
ground_truth = [None for _ in range(len(prompts))]

tokenizer = get_tokenizer()._tokenizer
model = get_model(functools.partial(model_provider, parallel_output=True), wrap_with_ddp=False)
model = get_model(functools.partial(model_provider, modelopt_gpt_mamba_builder), wrap_with_ddp=False)

report_current_memory_info()

Expand Down
Loading
Loading