preprocess the data via the vocab.txt from the `model_dir` path Args: cfg(modelscope.utils.config.ConfigDict) : model config model_dir (str): model path
(self, cfg, model_dir, mode, *args, **kwargs)
| 27 | """ |
| 28 | |
| 29 | def __init__(self, cfg, model_dir, mode, *args, **kwargs): |
| 30 | """preprocess the data via the vocab.txt from the `model_dir` path |
| 31 | |
| 32 | Args: |
| 33 | cfg(modelscope.utils.config.ConfigDict) : model config |
| 34 | model_dir (str): model path |
| 35 | """ |
| 36 | self.cfg = cfg |
| 37 | self.mode = mode |
| 38 | self.language = self.cfg.model.get('language', 'en') |
| 39 | if os.path.exists(model_dir): |
| 40 | model_dir = os.path.abspath(model_dir) |
| 41 | if self.language == 'en': |
| 42 | tokenizer = OFATokenizer.from_pretrained(model_dir) |
| 43 | elif self.language in ['zh', 'cn']: |
| 44 | tokenizer = OFATokenizerZH.from_pretrained(model_dir) |
| 45 | else: |
| 46 | raise NotImplementedError |
| 47 | # there is some diff between here and our ofa code, |
| 48 | # there will be no need to use param: use_bpe |
| 49 | tokenizer.add_tokens(['<code_{}>'.format(i) for i in range(8192)]) |
| 50 | tokenizer.add_tokens(['<bin_{}>'.format(i) for i in range(1000)]) |
| 51 | if self.cfg.model.get('multimodal_type', 'default') == 'text2sql': |
| 52 | tokenizer.add_tokens(['>=', '<=']) |
| 53 | self.tokenizer = tokenizer |
| 54 | self.bos_item = torch.LongTensor([tokenizer.bos_token_id]) |
| 55 | self.pad_item = torch.LongTensor([tokenizer.pad_token_id]) |
| 56 | self.eos_item = torch.LongTensor([tokenizer.eos_token_id]) |
| 57 | self.tgt_dict = self.src_dict = { |
| 58 | value: key |
| 59 | for key, value in tokenizer.get_vocab().items() |
| 60 | } |
| 61 | self.max_src_length = cfg.model.get('max_src_length', 256) |
| 62 | self.max_tgt_length = cfg.model.get('max_tgt_length', 256) |
| 63 | self.max_image_size = cfg.model.get('max_image_size', 512) |
| 64 | self.language = self.cfg.model.get('language', 'en') |
| 65 | self.prompt_type = self.cfg.model.get('prompt_type', 'none') |
| 66 | seed = self.cfg.model.get('seed', 7) |
| 67 | np.random.seed(seed) |
| 68 | set_torch_seed(seed) |
| 69 | imagenet_default_mean_and_std = self.cfg.model.get( |
| 70 | 'imagenet_default_mean_and_std', False) |
| 71 | if imagenet_default_mean_and_std: |
| 72 | self.mean = [0.485, 0.456, 0.406] |
| 73 | self.std = [0.229, 0.224, 0.225] |
| 74 | else: |
| 75 | self.mean = [0.5, 0.5, 0.5] |
| 76 | self.std = [0.5, 0.5, 0.5] |
| 77 | self.patch_image_size = self.cfg.model.get('patch_image_size', 480) |
| 78 | self.column_map = { |
| 79 | key: key |
| 80 | for key in OFA_TASK_KEY_MAPPING[self.cfg.task] |
| 81 | } |
| 82 | if hasattr(self.cfg, |
| 83 | 'dataset') and self.cfg.dataset.column_map is not None: |
| 84 | for k, v in self.cfg.dataset.column_map.items(): |
| 85 | self.column_map[k] = v |
| 86 | self.transtab = str.maketrans( |
nothing calls this directly
no test coverage detected