| 58 | from gsplat.compression_simulation.entropy_model import Entropy_factorized_optimized_refactor, Entropy_gaussian |
| 59 | |
| 60 | class ProfilerConfig: |
| 61 | def __init__(self): |
| 62 | self.enabled = False |
| 63 | self.activities = [ |
| 64 | torch.profiler.ProfilerActivity.CPU, |
| 65 | torch.profiler.ProfilerActivity.CUDA, |
| 66 | ] |
| 67 | |
| 68 | self.wait = 1 |
| 69 | self.warmup = 2 |
| 70 | self.active = 30_000 |
| 71 | |
| 72 | self.schedule = self._create_schedule() |
| 73 | |
| 74 | self.on_trace_ready = torch.profiler.tensorboard_trace_handler('./log/profiler') |
| 75 | self.record_shapes = True |
| 76 | self.profile_memory = True |
| 77 | self.with_stack = True |
| 78 | |
| 79 | def _create_schedule(self): |
| 80 | return torch.profiler.schedule( |
| 81 | wait=self.wait, |
| 82 | warmup=self.warmup, |
| 83 | active=self.active, |
| 84 | ) |
| 85 | |
| 86 | def update_schedule(self, **kwargs): |
| 87 | for key, value in kwargs.items(): |
| 88 | if hasattr(self, key): |
| 89 | setattr(self, key, value) |
| 90 | self.schedule = self._create_schedule() |
| 91 | |
| 92 | @dataclass |
| 93 | class CodecConfig: |