MCPcopy Create free account
hub / github.com/pytorch/pytorch / _build_model

Method _build_model

caffe2/python/models/seq2seq/train.py:101–151  ·  view source on GitHub ↗
(
        self,
        init_params,
    )

Source from the content-addressed store, hash-verified

99class Seq2SeqModelCaffe2:
100
101 def _build_model(
102 self,
103 init_params,
104 ):
105 model = Seq2SeqModelHelper(init_params=init_params)
106 self._build_shared(model)
107 self._build_embeddings(model)
108
109 forward_model = Seq2SeqModelHelper(init_params=init_params)
110 self._build_shared(forward_model)
111 self._build_embeddings(forward_model)
112
113 if self.num_gpus == 0:
114 loss_blobs = self.model_build_fun(model)
115 model.AddGradientOperators(loss_blobs)
116 self.norm_clipped_grad_update(
117 model,
118 scope='norm_clipped_grad_update'
119 )
120 self.forward_model_build_fun(forward_model)
121
122 else:
123 assert (self.batch_size % self.num_gpus) == 0
124
125 data_parallel_model.Parallelize_GPU(
126 forward_model,
127 input_builder_fun=lambda m: None,
128 forward_pass_builder_fun=self.forward_model_build_fun,
129 param_update_builder_fun=None,
130 devices=list(range(self.num_gpus)),
131 )
132
133 def clipped_grad_update_bound(model):
134 self.norm_clipped_grad_update(
135 model,
136 scope='norm_clipped_grad_update',
137 )
138
139 data_parallel_model.Parallelize_GPU(
140 model,
141 input_builder_fun=lambda m: None,
142 forward_pass_builder_fun=self.model_build_fun,
143 param_update_builder_fun=clipped_grad_update_bound,
144 devices=list(range(self.num_gpus)),
145 )
146 self.norm_clipped_sparse_grad_update(
147 model,
148 scope='norm_clipped_sparse_grad_update',
149 )
150 self.model = model
151 self.forward_net = forward_model.net
152
153 def _build_shared(self, model):
154 optimizer_params = self.model_params['optimizer_params']

Callers 1

Calls 10

_build_sharedMethod · 0.95
_build_embeddingsMethod · 0.95
model_build_funMethod · 0.95
Seq2SeqModelHelperClass · 0.90
listFunction · 0.85
rangeFunction · 0.50
AddGradientOperatorsMethod · 0.45

Tested by

no test coverage detected