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
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from megatron.training import print_rank_0
from megatron.training.checkpointing import load_checkpoint
from megatron.core import mpu
from megatron.training.arguments import parse_and_validate_args
from megatron.training.initialize import initialize_megatron
from megatron.legacy.model import GPTModel
from megatron.training import get_model
Expand Down Expand Up @@ -232,11 +233,12 @@ def generate_and_write_samples_conditional(model):
def main():
"""Main program."""

initialize_megatron(extra_args_provider=add_text_generate_args,
args_defaults={'tokenizer_type': 'GPT2BPETokenizer',
'no_load_rng': True,
'no_load_optim': True,
'seq_length': 2048})
parse_and_validate_args(extra_args_provider=add_text_generate_args,
args_defaults={'tokenizer_type': 'GPT2BPETokenizer',
'no_load_rng': True,
'no_load_optim': True,
'seq_length': 2048})
initialize_megatron()

# Set up model and load checkpoint
model = get_model(model_provider, wrap_with_ddp=False)
Expand Down
4 changes: 3 additions & 1 deletion examples/multimodal/model_converter/vision_model_tester.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from examples.multimodal.model import model_provider
from examples.multimodal.multimodal_args import add_multimodal_extra_args
from megatron.training import get_model
from megatron.training.arguments import parse_and_validate_args
from megatron.training.checkpointing import load_checkpoint
from megatron.training.initialize import initialize_megatron

Expand Down Expand Up @@ -50,7 +51,8 @@ def run_mcore_vision(model_path):
f"--pretrained-checkpoint={model_path}",
]

initialize_megatron(extra_args_provider=add_multimodal_extra_args)
parse_and_validate_args(extra_args_provider=add_multimodal_extra_args)
initialize_megatron()

def wrapped_model_provider(pre_process, post_process):
return model_provider(pre_process, post_process, parallel_output=False)
Expand Down
4 changes: 3 additions & 1 deletion examples/multimodal/run_text_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
VLMInferenceWrapper,
)
from megatron.training import get_args, get_model, get_tokenizer, print_rank_0, is_last_rank
from megatron.training.arguments import parse_and_validate_args
from megatron.training.checkpointing import load_checkpoint
from megatron.training.initialize import initialize_megatron

Expand Down Expand Up @@ -842,7 +843,8 @@ def run_evaluation_loop(model, configs, output_dir_override=None, iteration=None

def eval_tasks():
"""Vision language model text generation for single or batch tasks."""
initialize_megatron(extra_args_provider=add_text_generation_args)
parse_and_validate_args(extra_args_provider=add_text_generation_args)
initialize_megatron()

args = get_args()

Expand Down
4 changes: 3 additions & 1 deletion tools/run_dynamic_text_generation_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from megatron.inference.utils import add_inference_args, get_dynamic_inference_engine
from megatron.post_training.arguments import add_modelopt_args
from megatron.training import get_args
from megatron.training.arguments import parse_and_validate_args
from megatron.training.initialize import initialize_megatron


Expand Down Expand Up @@ -76,10 +77,11 @@ async def run_text_generation_server(

if __name__ == "__main__":
with torch.inference_mode():
initialize_megatron(
parse_and_validate_args(
extra_args_provider=add_text_generation_server_args,
args_defaults={'no_load_rng': True, 'no_load_optim': True},
)
initialize_megatron()

# Enable return_log_probs to allow prompt logprobs computation for echo=True requests
# This sets materialize_only_last_token_logits=False in the inference context,
Expand Down
4 changes: 3 additions & 1 deletion tools/run_inference_performance_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
from megatron.core import mpu
from megatron.training import get_args, get_model, get_tokenizer
from megatron.training.checkpointing import load_checkpoint
from megatron.training.arguments import parse_and_validate_args
from megatron.training.initialize import initialize_megatron

REQUEST_ID = 0
Expand Down Expand Up @@ -151,7 +152,7 @@ def generate_dynamic(
def main():
"""Main program."""

initialize_megatron(
parse_and_validate_args(
extra_args_provider=add_inference_benchmarking_args,
args_defaults={
'no_load_rng': True,
Expand All @@ -160,6 +161,7 @@ def main():
'exit_on_missing_checkpoint': True,
},
)
initialize_megatron()

args = get_args()

Expand Down
4 changes: 3 additions & 1 deletion tools/run_text_generation_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
from megatron.core import mpu
from megatron.training import get_args, get_model, get_tokenizer
from megatron.training.checkpointing import load_checkpoint
from megatron.training.arguments import parse_and_validate_args
from megatron.training.initialize import initialize_megatron


Expand Down Expand Up @@ -113,14 +114,15 @@ def add_text_generate_args(parser):
@torch.inference_mode()
def main(model_type: str = "gpt"):
"""Runs the text generation server with the specified model type."""
initialize_megatron(
parse_and_validate_args(
extra_args_provider=add_text_generate_args,
args_defaults={
'no_load_rng': True,
'no_load_optim': True,
'exit_on_missing_checkpoint': True,
},
)
initialize_megatron()
args = get_args()
if args.num_layers_per_virtual_pipeline_stage is not None:
print("Interleaved pipeline schedule is not yet supported for text generation.")
Expand Down
4 changes: 3 additions & 1 deletion tools/run_vlm_text_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from megatron.inference.text_generation.api import generate_and_post_process
from megatron.inference.text_generation.forward_step import ForwardStep
from megatron.training import get_args, get_model, print_rank_0
from megatron.training.arguments import parse_and_validate_args
from megatron.training.checkpointing import load_checkpoint
from megatron.training.initialize import initialize_megatron
from pretrain_vlm import model_provider
Expand Down Expand Up @@ -199,7 +200,8 @@ def main():

logging.getLogger(__name__).warning("Models using pipeline parallelism are not supported yet.")

initialize_megatron(extra_args_provider=add_text_generation_args)
parse_and_validate_args(extra_args_provider=add_text_generation_args)
initialize_megatron()

# Set up model and load checkpoint.
model = get_model(model_provider, wrap_with_ddp=False)
Expand Down
Loading