MCPcopy Create free account
hub / github.com/baidu/DDParser / train

Function train

ddparser/run.py:67–164  ·  view source on GitHub ↗

Train

(env)

Source from the content-addressed store, hash-verified

65
66
67def train(env):
68 """Train"""
69 args = env.args
70
71 logging.info("loading data.")
72 train = Corpus.load(args.train_data_path, env.fields)
73 dev = Corpus.load(args.valid_data_path, env.fields)
74 test = Corpus.load(args.test_data_path, env.fields)
75 logging.info("init dataset.")
76 train = TextDataset(train, env.fields, args.buckets)
77 dev = TextDataset(dev, env.fields, args.buckets)
78 test = TextDataset(test, env.fields, args.buckets)
79 logging.info("set the data loaders.")
80 train.loader = batchify(train, args.batch_size, args.use_data_parallel, True)
81 dev.loader = batchify(dev, args.batch_size)
82 test.loader = batchify(test, args.batch_size)
83
84 logging.info("{:6} {:5} sentences, ".format('train:', len(train)) + "{:3} batches, ".format(len(train.loader)) +
85 "{} buckets".format(len(train.buckets)))
86 logging.info("{:6} {:5} sentences, ".format('dev:', len(dev)) + "{:3} batches, ".format(len(dev.loader)) +
87 "{} buckets".format(len(dev.buckets)))
88 logging.info("{:6} {:5} sentences, ".format('test:', len(test)) + "{:3} batches, ".format(len(test.loader)) +
89 "{} buckets".format(len(test.buckets)))
90
91 logging.info("Create the model")
92 model = Model(args)
93
94 # init parallel strategy
95 if args.use_data_parallel:
96 dist.init_parallel_env()
97 model = paddle.DataParallel(model)
98
99 if args.encoding_model.startswith(
100 "ernie") and args.encoding_model != "ernie-lstm" or args.encoding_model == 'transformer':
101 args['lr'] = args.ernie_lr
102 else:
103 args['lr'] = args.lstm_lr
104
105 if args.encoding_model.startswith("ernie") and args.encoding_model != "ernie-lstm":
106 max_steps = 100 * len(train.loader)
107 decay = LinearDecay(args.lr, int(args.warmup_proportion * max_steps), max_steps)
108 else:
109 decay = dygraph.ExponentialDecay(learning_rate=args.lr, decay_steps=args.decay_steps, decay_rate=args.decay)
110
111 grad_clip = paddle.nn.ClipGradByGlobalNorm(clip_norm=args.clip)
112
113 if args.encoding_model.startswith("ernie") and args.encoding_model != "ernie-lstm":
114 optimizer = AdamW(
115 learning_rate=decay,
116 parameter_list=model.parameters(),
117 weight_decay=args.weight_decay,
118 grad_clip=grad_clip,
119 )
120 else:
121 optimizer = fluid.optimizer.AdamOptimizer(
122 learning_rate=decay,
123 beta1=args.mu,
124 beta2=args.nu,

Callers 1

run.pyFile · 0.70

Calls 11

TextDatasetClass · 0.90
batchifyFunction · 0.90
ModelClass · 0.90
LinearDecayClass · 0.90
AdamWClass · 0.90
MetricClass · 0.90
epoch_trainFunction · 0.90
epoch_evaluateFunction · 0.90
saveFunction · 0.90
loadFunction · 0.90
loadMethod · 0.45

Tested by

no test coverage detected