(self)
| 975 | return self._config.data_efficiency_config[DATA_ROUTING][RANDOM_LTD] |
| 976 | |
| 977 | def random_ltd_initialize(self): |
| 978 | assert self.random_ltd_enabled() |
| 979 | random_ltd_config = self.random_ltd_config() |
| 980 | random_ltd_queue = deque([x for x in sorted(random_ltd_config[RANDOM_LTD_LAYER_ID])]) |
| 981 | count = 0 |
| 982 | for name, layer in self.module.named_modules(): |
| 983 | if isinstance(layer, RandomLayerTokenDrop): |
| 984 | if len(random_ltd_queue) != 0 and str(random_ltd_queue[0]) in name: ###[1,2,3] |
| 985 | layer.init_config(random_ltd_config, self.random_ltd_scheduler, count) |
| 986 | random_ltd_queue.popleft() |
| 987 | count += 1 |
| 988 | |
| 989 | if random_ltd_config[RANDOM_LTD_LAYER_NUM] != count: |
| 990 | raise ValueError(f'random_ltd_layer_num {random_ltd_config[RANDOM_LTD_LAYER_NUM]} must be \ |
| 991 | equivalent to the len of random_ltd_layer_id {count}') |
| 992 | |
| 993 | if random_ltd_config[RANDOM_LTD_LAYER_TOKEN_LR_SCHEDULE][RANDOM_LTD_LAYER_TOKEN_LR_ENABLED]: |
| 994 | assert self.client_lr_scheduler is None |
| 995 | raise ValueError('not yet support') |
| 996 | #self.lr_scheduler = lr_schedules.WarmupLayerTokenDecayLR(self.optimizer, self.random_ltd_scheduler) |
| 997 | |
| 998 | def get_data_parallel_rank(self): |
| 999 | return groups.get_data_parallel_rank() |
no test coverage detected