A context container which contains the running states, such as the layer structures, log folder, hooks to run, is pre_forward or after forward, etc. Arguments: dest_folder: str The working directory set to debug context to set where the hook dumped data/info.
| 11 | |
| 12 | |
| 13 | class DebuggerContext: |
| 14 | """ |
| 15 | A context container which contains the running states, such as the layer structures, |
| 16 | log folder, hooks to run, is pre_forward or after forward, etc. |
| 17 | |
| 18 | Arguments: |
| 19 | dest_folder: str |
| 20 | The working directory set to debug context to set where the hook dumped data/info. |
| 21 | """ |
| 22 | |
| 23 | def __init__(self, dest_folder: str = None): |
| 24 | self.pre_forward_actions = [] |
| 25 | self.after_forward_actions = [] |
| 26 | |
| 27 | self.layer_names = [] |
| 28 | self.layer_inner_counter = [] |
| 29 | |
| 30 | self.module_forward_hook_handle = None |
| 31 | self.module_forward_pre_hook_handle = None |
| 32 | |
| 33 | self.forward_hook_handles = {} # module to handlers |
| 34 | self.forward_pre_hook_handles = {} |
| 35 | self.log_folder = dest_folder |
| 36 | self.is_pre_forward = True |
| 37 | self._init_log_folder() |
| 38 | |
| 39 | def _init_log_folder(self): |
| 40 | if self.log_folder is None: |
| 41 | pwd = os.getcwd() |
| 42 | self.log_folder = os.path.join(pwd, "data_dump") |
| 43 | |
| 44 | rank = tensorrt_llm.mpi_rank() |
| 45 | |
| 46 | p = Path(self.log_folder) / f"rank{rank}" |
| 47 | self.log_folder = p.absolute() |
| 48 | p.mkdir(parents=True, exist_ok=True) |
| 49 | |
| 50 | def get_log_folder(self): |
| 51 | return self.log_folder |
| 52 | |
| 53 | def check_in_pre_forward(self): |
| 54 | return self.is_pre_forward |
| 55 | |
| 56 | def mark_in_pre_forward(self, is_pre_forward): |
| 57 | self.is_pre_forward = is_pre_forward |
| 58 | |
| 59 | def clear_state(self): |
| 60 | self.pre_forward_actions.clear() |
| 61 | self.after_forward_actions.clear() |
| 62 | self.layer_names.clear() |
| 63 | self.layer_inner_counter.clear() |
| 64 | |
| 65 | if self.module_forward_hook_handle is not None: |
| 66 | self.module_forward_hook_handle.remove() |
| 67 | if self.module_forward_pre_hook_handle is not None: |
| 68 | self.module_forward_pre_hook_handle.remove() |
| 69 | |
| 70 | self.module_forward_hook_handle = None |