(self)
| 1029 | return self._config.data_efficiency_config[DATA_ROUTING][RANDOM_LTD] |
| 1030 | |
| 1031 | def random_ltd_initialize(self): |
| 1032 | assert self.random_ltd_enabled() |
| 1033 | random_ltd_config = self.random_ltd_config() |
| 1034 | random_ltd_queue = deque([x for x in sorted(random_ltd_config[RANDOM_LTD_LAYER_ID])]) |
| 1035 | count = 0 |
| 1036 | for name, layer in self.module.named_modules(): |
| 1037 | if isinstance(layer, RandomLayerTokenDrop): |
| 1038 | if len(random_ltd_queue) != 0 and str(random_ltd_queue[0]) in name: ###[1,2,3] |
| 1039 | layer.init_config(random_ltd_config, self.random_ltd_scheduler, count) |
| 1040 | random_ltd_queue.popleft() |
| 1041 | count += 1 |
| 1042 | |
| 1043 | if random_ltd_config[RANDOM_LTD_LAYER_NUM] != count: |
| 1044 | raise ValueError(f'random_ltd_layer_num {random_ltd_config[RANDOM_LTD_LAYER_NUM]} must be \ |
| 1045 | equivalent to the len of random_ltd_layer_id {count}') |
| 1046 | |
| 1047 | if random_ltd_config[RANDOM_LTD_LAYER_TOKEN_LR_SCHEDULE][RANDOM_LTD_LAYER_TOKEN_LR_ENABLED]: |
| 1048 | assert self.client_lr_scheduler is None |
| 1049 | raise ValueError('not yet support') |
| 1050 | #self.lr_scheduler = lr_schedules.WarmupLayerTokenDecayLR(self.optimizer, self.random_ltd_scheduler) |
| 1051 | |
| 1052 | def get_data_parallel_rank(self): |
| 1053 | return groups.get_data_parallel_rank() |
no test coverage detected