MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / load_model

Function load_model

codegeex/mindspore/generation_values.py:39–171  ·  view source on GitHub ↗

r""" The main function for load model

(args_opt)

Source from the content-addressed store, hash-verified

37
38
39def load_model(args_opt):
40 r"""
41 The main function for load model
42 """
43 # Set execution mode
44 context.set_context(save_graphs=False,
45 mode=context.GRAPH_MODE,
46 device_target=args_opt.device_target)
47 context.set_context(variable_memory_max_size="30GB")
48 # Set parallel context
49 if args_opt.distribute == "true":
50 D.init()
51 device_num = D.get_group_size()
52 rank = D.get_rank()
53 print("rank_id is {}, device_num is {}".format(rank, device_num))
54 context.reset_auto_parallel_context()
55 context.set_auto_parallel_context(
56 parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL,
57 gradients_mean=False,
58 full_batch=True,
59 loss_repeated_mean=True,
60 enable_parallel_optimizer=False,
61 pipeline_stages=args_opt.stage_num)
62 set_algo_parameters(elementwise_op_strategy_follow=True)
63 _set_multi_subgraphs()
64
65 else:
66 rank = 0
67 device_num = 1
68 context.reset_auto_parallel_context()
69 context.set_auto_parallel_context(
70 strategy_ckpt_load_file=args_opt.strategy_load_ckpt_path)
71 context.set_context(
72 save_graphs=False,
73 save_graphs_path="/cache/graphs_of_device_id_" + str(rank),
74 )
75 use_past = (args_opt.use_past == "true")
76 print('local_rank:{}, start to run...'.format(rank), flush=True)
77 if args_opt.export:
78 use_past = True
79 # Set model property
80 model_parallel_num = args_opt.op_level_model_parallel_num
81 data_parallel_num = int(device_num / model_parallel_num)
82
83 parallel_config = TransformerOpParallelConfig(data_parallel=data_parallel_num,
84 model_parallel=model_parallel_num,
85 pipeline_stage=args_opt.stage_num,
86 micro_batch_num=args_opt.micro_size,
87 optimizer_shard=False,
88 vocab_emb_dp=bool(args_opt.word_emb_dp),
89 recompute=True)
90
91 per_batch_size = args_opt.per_batch_size
92 batch_size = per_batch_size * data_parallel_num
93 config = PanguAlphaConfig(
94 batch_size=batch_size,
95 seq_length=args_opt.seq_length,
96 vocab_size=args_opt.vocab_size,

Callers 1

mainFunction · 0.70

Calls 6

PanguAlphaConfigClass · 0.90
PanguAlphaModelClass · 0.90
LogitsNetClass · 0.90
load_checkpointFunction · 0.85
existsMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected