# Loading augmented train, validation and test sets BASE = r"augmented" def load_split(path: str) -> Dataset: with open(path, encoding="utf-8") as f: d = json.load(f) return Dataset.from_dict({"input_embeds": d["X"], "labels": d["y"]}) train_dataset = load_split(f"{BASE}/train.json") val_dataset = load_split(f"{BASE}/val.json") test_dataset = load_split(f"{BASE}/test.json") # Formulate embedding dimension EMBED_DIM = len(train_dataset[0]["input_embeds"]) # Return Pytorch tensors train_dataset.set_format("torch") val_dataset.set_format("torch") test_dataset.set_format("torch")