Diff Coverage

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

Source File Diff Coverage (%) Missing Lines
hyper_parallel/trainer/base.py 62.5% 715-716,729,734,737,739,741-742,744
hyper_parallel/trainer/runtime/data_iterator.py 86.8% 74,77,90-92
hyper_parallel/trainer/base.py
711
712
713
714
715
716
717
718
719
        try:
            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
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
                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}")

                if config.dataloader.use_background_prefetcher:
                    self.data_iterator.stop()

            self.on_train_end()
        finally:
            if config.dataloader.use_background_prefetcher:
                self.data_iterator.stop()
hyper_parallel/trainer/runtime/data_iterator.py
70
71
72
73
74
75
76
77
78
79
80
81
                # don't mutate the captured state in-place. The underlying dataloader's
                # state_dict() should handle deepcopying if necessary.
                state = self.original_state_dict() if self.original_state_dict else None
                if not self._put_result((item, state)):
                    break
        # The worker must transfer any producer failure back to the training thread.
        except Exception as exc:  # pylint: disable=broad-exception-caught
            self._put_result((exc, None))
        finally:
            # A timed-out stop cannot cancel next(), so the worker must release
            # its own references when that call eventually returns.
            if self.stop_event.is_set():
86
87
88
89
90
91
92
93
94
95
96
        while not self.stop_event.is_set():
            try:
                self.queue.put(result, timeout=0.1)
                return True
            except queue.Full:
                continue
        return False

    def _drain_queue(self) -> None:
        """Release every result currently retained by the prefetch queue."""
        while True: