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

Function load_model

codegeex/mindspore/generation_humaneval.py:42–174  ·  view source on GitHub ↗

r""" The main function for load model

(args_opt)

Source from the content-addressed store, hash-verified

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

Callers 1

mainFunction · 0.70

Calls 6

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

Tested by

no test coverage detected