Diff Coverage

Diff: origin/master...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/auto_parallel/sapp_nd/nd/common/cost_model_preprocess.py 100%  
hyper_parallel/auto_parallel/sapp_nd/nd/common/framework_parsers/cost_model_parser_hyper.py 100%  
hyper_parallel/auto_parallel/sapp_nd/nd/common/framework_parsers/cost_model_parser_mindformers.py 100%  
hyper_parallel/components/checkpoint/registry.py 85.7% 71
hyper_parallel/trainer/base.py 73.0% 696,717-718,731,736,738-739,741,743-744
hyper_parallel/trainer/callbacks/base.py 100%  
hyper_parallel/components/checkpoint/registry.py
67
68
69
70
71
72
73
74
75

    def __delitem__(self, key: str) -> None:
        """Delete every registration for ``key`` from this registry."""
        if key not in self._local_mapping and key not in self._global_mapping:
            raise KeyError(key)
        self._local_mapping.pop(key, None)
        self._global_mapping.pop(key, None)

    def __iter__(self) -> Iterator[str]:
hyper_parallel/trainer/base.py
692
693
694
695
696
697
698
699
700

        try:
            empty_cache()
            dist.barrier()
            synchronize()
        finally:
            destroy_process_group()

    def train(self) -> None:
713
714
715
716
717
718
719
720
721

            start_epoch = self.state.epoch
            for epoch in range(start_epoch, self.train_epochs):
                if epoch != start_epoch:
                    self.train_dataloader.set_epoch(epoch)
                    self.data_iterator = HyperIter(
                        self.train_dataloader, use_background_prefetcher=config.dataloader.use_background_prefetcher
                    )
                self.state.epoch = epoch
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
                for _ in range(start_step, train_steps):
                    try:
                        self.train_step(self.data_iterator)
                    except StopIteration:
                        logger.info(
                            "epoch:%s Dataloader finished with drop_last %s",
                            epoch,
                            config.dataloader.drop_last,
                        )
                        break

                self.on_epoch_end()
                self.state.epoch = epoch + 1

                print_device_mem_info(f"VRAM usage after epoch {epoch + 1}")

                self.data_iterator.stop()
                self.data_iterator = None
        finally:
            # ExitStack runs every callback even when an earlier cleanup raises.
            with ExitStack() as cleanup:
                cleanup.callback(self.destroy_distributed)