(model, blob_in, blob_out, dim_in,
RunningMeanInitializer=None, RunningVarianceInitializer=None,
order="NCHW", **kwargs)
| 276 | return normalized, mean, std |
| 277 | |
| 278 | def moments_with_running_stats(model, blob_in, blob_out, dim_in, |
| 279 | RunningMeanInitializer=None, RunningVarianceInitializer=None, |
| 280 | order="NCHW", **kwargs): |
| 281 | |
| 282 | if model.init_params: |
| 283 | rm_init = ("ConstantFill", {'value': 0.0}) |
| 284 | riv_init = ("ConstantFill", {'value': 1.0}) |
| 285 | |
| 286 | RunningMeanInitializer = initializers.update_initializer( |
| 287 | RunningMeanInitializer, rm_init, ("ConstantFill", {}) |
| 288 | ) |
| 289 | RunningVarianceInitializer = initializers.update_initializer( |
| 290 | RunningVarianceInitializer, riv_init, ("ConstantFill", {}) |
| 291 | ) |
| 292 | else: |
| 293 | RunningMeanInitializer = initializers.ExternalInitializer() |
| 294 | RunningVarianceInitializer = initializers.ExternalInitializer() |
| 295 | |
| 296 | running_mean = model.create_param( |
| 297 | param_name=blob_out + '_rm', |
| 298 | shape=[dim_in], |
| 299 | initializer=RunningMeanInitializer, |
| 300 | tags=ParameterTags.COMPUTED_PARAM |
| 301 | ) |
| 302 | |
| 303 | # this is just running variance |
| 304 | running_inv_var = model.create_param( |
| 305 | param_name=blob_out + '_riv', |
| 306 | shape=[dim_in], |
| 307 | initializer=RunningVarianceInitializer, |
| 308 | tags=ParameterTags.COMPUTED_PARAM |
| 309 | ) |
| 310 | |
| 311 | blob_outs = [blob_out + "_sm", blob_out + "_sv"] |
| 312 | if order == 'NCHW': |
| 313 | blob_outputs = model.net.Moments( |
| 314 | [blob_in], blob_outs, |
| 315 | axes=[0, 2, 3], |
| 316 | order=order, keepdims=False, **kwargs) |
| 317 | elif order == 'NHWC': |
| 318 | blob_outputs = model.net.Moments( |
| 319 | [blob_in], blob_outs, |
| 320 | axes=[0, 1, 2], |
| 321 | order=order, keepdims=False, **kwargs) |
| 322 | return blob_outputs |
nothing calls this directly
no test coverage detected
searching dependent graphs…