(self)
| 917 | return self._config.data_efficiency_config[DATA_ROUTING][RANDOM_LTD] |
| 918 | |
| 919 | def random_ltd_initialize(self): |
| 920 | assert self.random_ltd_enabled() |
| 921 | random_ltd_config = self.random_ltd_config() |
| 922 | random_ltd_queue = deque([x for x in sorted(random_ltd_config[RANDOM_LTD_LAYER_ID])]) |
| 923 | count = 0 |
| 924 | for name, layer in self.module.named_modules(): |
| 925 | if isinstance(layer, RandomLayerTokenDrop): |
| 926 | if len(random_ltd_queue) != 0 and str(random_ltd_queue[0]) in name: ###[1,2,3] |
| 927 | layer.init_config(random_ltd_config, self.random_ltd_scheduler, count) |
| 928 | random_ltd_queue.popleft() |
| 929 | count += 1 |
| 930 | |
| 931 | if random_ltd_config[RANDOM_LTD_LAYER_NUM] != count: |
| 932 | raise ValueError(f'random_ltd_layer_num {random_ltd_config[RANDOM_LTD_LAYER_NUM]} must be \ |
| 933 | equivalent to the len of random_ltd_layer_id {count}') |
| 934 | |
| 935 | if random_ltd_config[RANDOM_LTD_LAYER_TOKEN_LR_SCHEDULE][RANDOM_LTD_LAYER_TOKEN_LR_ENABLED]: |
| 936 | assert self.client_lr_scheduler is None |
| 937 | raise ValueError('not yet support') |
| 938 | #self.lr_scheduler = lr_schedules.WarmupLayerTokenDecayLR(self.optimizer, self.random_ltd_scheduler) |
| 939 | |
| 940 | def get_data_parallel_rank(self): |
| 941 | return groups.get_data_parallel_rank() |
no test coverage detected