Skip to content

Commit

Permalink
modify api in unit test
Browse files Browse the repository at this point in the history
  • Loading branch information
pkuzyc committed Sep 20, 2023
1 parent f408ef7 commit 3e885f3
Show file tree
Hide file tree
Showing 2 changed files with 80 additions and 34 deletions.
3 changes: 0 additions & 3 deletions paddle/fluid/distributed/auto_parallel/spmd_rules/rules.h
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,6 @@ namespace paddle {
namespace distributed {
namespace auto_parallel {

// layer_norm rule
REGISTER_SPMD_RULE(layer_norm, LayerNormSPMDRule);

// replicated rule
REGISTER_SPMD_RULE(replicated, ReplicatedSPMDRule);

Expand Down
111 changes: 80 additions & 31 deletions test/auto_parallel/spmd_rules/test_layer_norm_rule.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,8 +64,11 @@ def test_infer_forward(self):
self.scale_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_forward(
[self.x_spec, self.scale_spec, self.bias_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand All @@ -90,8 +93,11 @@ def test_infer_forward(self):
self.bias_spec.set_dims_mapping([0])

result_dist_attrs = self.rule.infer_forward(
[self.x_spec, self.scale_spec, self.bias_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -120,8 +126,11 @@ def test_infer_forward(self):
self.bias_spec.set_dims_mapping([1])

result_dist_attrs = self.rule.infer_forward(
[self.x_spec, self.scale_spec, self.bias_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -158,9 +167,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -198,9 +212,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([0])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -238,9 +257,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -278,9 +302,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -317,11 +346,16 @@ def test_infer_backward(self):
self.mean_spec.set_dims_mapping([0])
self.var_spec.set_dims_mapping([-1])

with self.assertRaises(BaseException):
with self.assertRaises(NotImplementedError):
result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)

# [-1, 1, -1], [0], [-1] (outputs) -->
Expand All @@ -346,9 +380,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -386,9 +425,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down Expand Up @@ -426,9 +470,14 @@ def test_infer_backward(self):
self.var_spec.set_dims_mapping([-1])

result_dist_attrs = self.rule.infer_backward(
[self.x_spec, self.scale_spec, self.bias_spec],
[self.out_spec, self.mean_spec, self.var_spec],
list(self.attrs.values()),
self.x_spec,
self.scale_spec,
self.bias_spec,
self.out_spec,
self.mean_spec,
self.var_spec,
self.attrs['epsilon'],
self.attrs['begin_norm_axis'],
)
infered_input_dist_attrs = result_dist_attrs[0]
infered_output_dist_attrs = result_dist_attrs[1]
Expand Down

0 comments on commit 3e885f3

Please sign in to comment.