|
12 | 12 | # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. |
13 | 13 | # See the License for the specific language governing permissions and |
14 | 14 | # limitations under the License. |
15 | | -import pdb |
16 | 15 | import tempfile |
17 | 16 | import unittest |
18 | 17 |
|
@@ -618,3 +617,81 @@ def test_full_loop_no_noise(self): |
618 | 617 |
|
619 | 618 | assert abs(result_sum.item() - 10629923278.7104) < 1e-2 |
620 | 619 | assert abs(result_mean.item() - 13841045.9358) < 1e-3 |
| 620 | + |
| 621 | + def test_from_pretrained_save_pretrained(self): |
| 622 | + kwargs = dict(self.forward_default_kwargs) |
| 623 | + |
| 624 | + num_inference_steps = kwargs.pop("num_inference_steps", None) |
| 625 | + |
| 626 | + for scheduler_class in self.scheduler_classes: |
| 627 | + sample = self.dummy_sample |
| 628 | + residual = 0.1 * sample |
| 629 | + |
| 630 | + scheduler_config = self.get_scheduler_config() |
| 631 | + scheduler = scheduler_class(**scheduler_config) |
| 632 | + |
| 633 | + with tempfile.TemporaryDirectory() as tmpdirname: |
| 634 | + scheduler.save_config(tmpdirname) |
| 635 | + new_scheduler = scheduler_class.from_config(tmpdirname) |
| 636 | + |
| 637 | + if num_inference_steps is not None and hasattr(scheduler, "set_timesteps"): |
| 638 | + scheduler.set_timesteps(num_inference_steps) |
| 639 | + new_scheduler.set_timesteps(num_inference_steps) |
| 640 | + elif num_inference_steps is not None and not hasattr(scheduler, "set_timesteps"): |
| 641 | + kwargs["num_inference_steps"] = num_inference_steps |
| 642 | + |
| 643 | + output = scheduler.step_pred(residual, 1, sample, **kwargs)["prev_sample"] |
| 644 | + new_output = new_scheduler.step_pred(residual, 1, sample, **kwargs)["prev_sample"] |
| 645 | + |
| 646 | + assert np.sum(np.abs(output - new_output)) < 1e-5, "Scheduler outputs are not identical" |
| 647 | + |
| 648 | + def test_step_shape(self): |
| 649 | + kwargs = dict(self.forward_default_kwargs) |
| 650 | + |
| 651 | + num_inference_steps = kwargs.pop("num_inference_steps", None) |
| 652 | + |
| 653 | + for scheduler_class in self.scheduler_classes: |
| 654 | + scheduler_config = self.get_scheduler_config() |
| 655 | + scheduler = scheduler_class(**scheduler_config) |
| 656 | + |
| 657 | + sample = self.dummy_sample |
| 658 | + residual = 0.1 * sample |
| 659 | + |
| 660 | + if num_inference_steps is not None and hasattr(scheduler, "set_timesteps"): |
| 661 | + scheduler.set_timesteps(num_inference_steps) |
| 662 | + elif num_inference_steps is not None and not hasattr(scheduler, "set_timesteps"): |
| 663 | + kwargs["num_inference_steps"] = num_inference_steps |
| 664 | + |
| 665 | + output_0 = scheduler.step_pred(residual, 0, sample, **kwargs)["prev_sample"] |
| 666 | + output_1 = scheduler.step_pred(residual, 1, sample, **kwargs)["prev_sample"] |
| 667 | + |
| 668 | + self.assertEqual(output_0.shape, sample.shape) |
| 669 | + self.assertEqual(output_0.shape, output_1.shape) |
| 670 | + |
| 671 | + def test_pytorch_equal_numpy(self): |
| 672 | + kwargs = dict(self.forward_default_kwargs) |
| 673 | + |
| 674 | + num_inference_steps = kwargs.pop("num_inference_steps", None) |
| 675 | + |
| 676 | + for scheduler_class in self.scheduler_classes: |
| 677 | + sample = self.dummy_sample |
| 678 | + residual = 0.1 * sample |
| 679 | + |
| 680 | + sample_pt = torch.tensor(sample) |
| 681 | + residual_pt = 0.1 * sample_pt |
| 682 | + |
| 683 | + scheduler_config = self.get_scheduler_config() |
| 684 | + scheduler = scheduler_class(**scheduler_config) |
| 685 | + |
| 686 | + scheduler_pt = scheduler_class(tensor_format="pt", **scheduler_config) |
| 687 | + |
| 688 | + if num_inference_steps is not None and hasattr(scheduler, "set_timesteps"): |
| 689 | + scheduler.set_timesteps(num_inference_steps) |
| 690 | + scheduler_pt.set_timesteps(num_inference_steps) |
| 691 | + elif num_inference_steps is not None and not hasattr(scheduler, "set_timesteps"): |
| 692 | + kwargs["num_inference_steps"] = num_inference_steps |
| 693 | + |
| 694 | + output = scheduler.step_pred(residual, 1, sample, **kwargs)["prev_sample"] |
| 695 | + output_pt = scheduler_pt.step_pred(residual_pt, 1, sample_pt, **kwargs)["prev_sample"] |
| 696 | + |
| 697 | + assert np.sum(np.abs(output - output_pt.numpy())) < 1e-4, "Scheduler outputs are not identical" |
0 commit comments