(
self,
tokenizer,
data_dir,
type_path="train",
max_source_length=1024,
max_target_length=56,
n_obs=None,
overwrite_cache=False,
prefix="",
)
| 78 | |
| 79 | class SummarizationDataset(Dataset): |
| 80 | def __init__( |
| 81 | self, |
| 82 | tokenizer, |
| 83 | data_dir, |
| 84 | type_path="train", |
| 85 | max_source_length=1024, |
| 86 | max_target_length=56, |
| 87 | n_obs=None, |
| 88 | overwrite_cache=False, |
| 89 | prefix="", |
| 90 | ): |
| 91 | super().__init__() |
| 92 | tok_name = tokenizer.__class__.__name__.lower().rstrip("tokenizer") |
| 93 | self.source = encode_file( |
| 94 | tokenizer, |
| 95 | os.path.join(data_dir, type_path + ".source"), |
| 96 | max_source_length, |
| 97 | overwrite_cache=overwrite_cache, |
| 98 | prefix=prefix, |
| 99 | tok_name=tok_name, |
| 100 | ) |
| 101 | tgt_path = os.path.join(data_dir, type_path + ".target") |
| 102 | if hasattr(tokenizer, "set_lang"): |
| 103 | tokenizer.set_lang("ro_RO") # HACK: only applies to mbart |
| 104 | self.target = encode_file( |
| 105 | tokenizer, tgt_path, max_target_length, overwrite_cache=overwrite_cache, tok_name=tok_name |
| 106 | ) |
| 107 | if n_obs is not None: |
| 108 | self.source = self.source[:n_obs] |
| 109 | self.target = self.target[:n_obs] |
| 110 | self.pad_token_id = tokenizer.pad_token_id |
| 111 | |
| 112 | def __len__(self): |
| 113 | return len(self.source) |
nothing calls this directly
no test coverage detected