Buckets:
tahamajs/Sysmem2_in_AI / ComputerAssignments /CA6_systematic_generalization /src /datasets /dataset_implementations.py
| import torch | |
| from torch.utils.data import Dataset, DataLoader | |
| from typing import List, Dict, Tuple, Any, Optional, Union | |
| import itertools | |
| import random | |
| import numpy as np | |
| from dataclasses import dataclass | |
| import json | |
| import os | |
| from ..core.components import ( | |
| SystematicGeneralizationTask, | |
| Component, | |
| ComponentType, | |
| CompositeExpression, | |
| ExpressionGenerator, | |
| ) | |
| class SCANDataset(Dataset): | |
| def __init__( | |
| self, split_type: str = "simple", max_length: int = 20, data_dir: str = None | |
| ): | |
| self.split_type = split_type | |
| self.max_length = max_length | |
| self.data_dir = data_dir or "data" | |
| self.actions = ["walk", "look", "run", "jump", "turn"] | |
| self.directions = ["left", "right", "around", "opposite"] | |
| self.modifiers = ["twice", "thrice", "after"] | |
| self.vocab = ( | |
| ["<PAD>", "<SOS>", "<EOS>"] | |
| + self.actions | |
| + self.directions | |
| + self.modifiers | |
| + ["and"] | |
| ) | |
| self.vocab_size = len(self.vocab) | |
| self.word_to_idx = {word: idx for idx, word in enumerate(self.vocab)} | |
| self.idx_to_word = {idx: word for idx, word in enumerate(self.vocab)} | |
| self.data = self._generate_scan_data() | |
| self._save_vocabulary() | |
| def _generate_scan_data(self) -> List[Dict[str, Any]]: | |
| data = [] | |
| for action in self.actions: | |
| for direction in self.directions: | |
| if direction in ["left", "right"]: | |
| command = f"{action} {direction}" | |
| if action == "walk": | |
| action_seq = ( | |
| ["I_TURN_LEFT", "I_WALK"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_WALK"] | |
| ) | |
| elif action == "look": | |
| action_seq = ( | |
| ["I_TURN_LEFT", "I_LOOK"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_LOOK"] | |
| ) | |
| elif action == "run": | |
| action_seq = ( | |
| ["I_TURN_LEFT", "I_RUN"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_RUN"] | |
| ) | |
| elif action == "jump": | |
| action_seq = ( | |
| ["I_TURN_LEFT", "I_JUMP"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_JUMP"] | |
| ) | |
| else: | |
| action_seq = ( | |
| ["I_TURN_LEFT"] if direction == "left" else ["I_TURN_RIGHT"] | |
| ) | |
| data.append( | |
| { | |
| "command": command, | |
| "actions": action_seq, | |
| "command_tokens": self._tokenize(command), | |
| "action_tokens": self._tokenize(" ".join(action_seq)), | |
| "complexity": 1, | |
| "is_compositional": False, | |
| } | |
| ) | |
| for action in self.actions[:3]: | |
| for direction in ["left", "right"]: | |
| for modifier in ["twice", "thrice"]: | |
| command = f"{action} {direction} {modifier}" | |
| if action == "walk": | |
| base_seq = ( | |
| ["I_TURN_LEFT", "I_WALK"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_WALK"] | |
| ) | |
| elif action == "look": | |
| base_seq = ( | |
| ["I_TURN_LEFT", "I_LOOK"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_LOOK"] | |
| ) | |
| else: | |
| base_seq = ( | |
| ["I_TURN_LEFT", "I_RUN"] | |
| if direction == "left" | |
| else ["I_TURN_RIGHT", "I_RUN"] | |
| ) | |
| if modifier == "twice": | |
| action_seq = base_seq + base_seq | |
| else: | |
| action_seq = base_seq + base_seq + base_seq | |
| data.append( | |
| { | |
| "command": command, | |
| "actions": action_seq, | |
| "command_tokens": self._tokenize(command), | |
| "action_tokens": self._tokenize(" ".join(action_seq)), | |
| "complexity": 2, | |
| "is_compositional": True, | |
| } | |
| ) | |
| for action1, action2 in itertools.product(self.actions[:2], self.actions[:2]): | |
| if action1 != action2: | |
| for direction1, direction2 in itertools.product( | |
| ["left", "right"], ["left", "right"] | |
| ): | |
| command = f"{action1} {direction1} and {action2} {direction2}" | |
| if action1 == "walk": | |
| seq1 = ( | |
| ["I_TURN_LEFT", "I_WALK"] | |
| if direction1 == "left" | |
| else ["I_TURN_RIGHT", "I_WALK"] | |
| ) | |
| else: | |
| seq1 = ( | |
| ["I_TURN_LEFT", "I_LOOK"] | |
| if direction1 == "left" | |
| else ["I_TURN_RIGHT", "I_LOOK"] | |
| ) | |
| if action2 == "walk": | |
| seq2 = ( | |
| ["I_TURN_LEFT", "I_WALK"] | |
| if direction2 == "left" | |
| else ["I_TURN_RIGHT", "I_WALK"] | |
| ) | |
| else: | |
| seq2 = ( | |
| ["I_TURN_LEFT", "I_LOOK"] | |
| if direction2 == "left" | |
| else ["I_TURN_RIGHT", "I_LOOK"] | |
| ) | |
| action_seq = seq1 + seq2 | |
| data.append( | |
| { | |
| "command": command, | |
| "actions": action_seq, | |
| "command_tokens": self._tokenize(command), | |
| "action_tokens": self._tokenize(" ".join(action_seq)), | |
| "complexity": 3, | |
| "is_compositional": True, | |
| } | |
| ) | |
| return data | |
| def _tokenize(self, text: str) -> List[int]: | |
| words = text.lower().split() | |
| tokens = [self.word_to_idx.get(word, 0) for word in words] | |
| return tokens | |
| def _save_vocabulary(self): | |
| os.makedirs(self.data_dir, exist_ok=True) | |
| vocab_path = os.path.join(self.data_dir, "scan_vocab.json") | |
| vocab_data = { | |
| "vocab": self.vocab, | |
| "word_to_idx": self.word_to_idx, | |
| "idx_to_word": self.idx_to_word, | |
| "vocab_size": self.vocab_size, | |
| } | |
| with open(vocab_path, "w") as f: | |
| json.dump(vocab_data, f, indent=2) | |
| def __len__(self) -> int: | |
| return len(self.data) | |
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: | |
| item = self.data[idx] | |
| command_tokens = item["command_tokens"][: self.max_length] | |
| command_tokens += [0] * (self.max_length - len(command_tokens)) | |
| action_tokens = item["action_tokens"][: self.max_length] | |
| action_tokens += [0] * (self.max_length - len(action_tokens)) | |
| return { | |
| "command": torch.tensor(command_tokens, dtype=torch.long), | |
| "actions": torch.tensor(action_tokens, dtype=torch.long), | |
| "command_text": item["command"], | |
| "action_text": " ".join(item["actions"]), | |
| "complexity": item["complexity"], | |
| "is_compositional": item["is_compositional"], | |
| } | |
| def get_statistics(self) -> Dict[str, Any]: | |
| complexities = [item["complexity"] for item in self.data] | |
| compositional_count = sum(1 for item in self.data if item["is_compositional"]) | |
| return { | |
| "total_examples": len(self.data), | |
| "compositional_examples": compositional_count, | |
| "non_compositional_examples": len(self.data) - compositional_count, | |
| "vocab_size": self.vocab_size, | |
| "avg_command_length": np.mean( | |
| [len(item["command"].split()) for item in self.data] | |
| ), | |
| "avg_action_length": np.mean([len(item["actions"]) for item in self.data]), | |
| "complexity_distribution": { | |
| "1": complexities.count(1), | |
| "2": complexities.count(2), | |
| "3": complexities.count(3), | |
| }, | |
| } | |
| class ArithmeticReasoningDataset(Dataset): | |
| def __init__( | |
| self, | |
| number_range: Tuple[int, int] = (1, 100), | |
| expression_length: Tuple[int, int] = (3, 15), | |
| data_dir: str = None, | |
| ): | |
| self.number_range = number_range | |
| self.expression_length = expression_length | |
| self.data_dir = data_dir or "data" | |
| self.operators = ["+", "-", "*", "/", "(", ")"] | |
| self.numbers = [str(i) for i in range(number_range[0], number_range[1] + 1)] | |
| self.vocab = ["<PAD>", "<SOS>", "<EOS>", "="] + self.numbers + self.operators | |
| self.vocab_size = len(self.vocab) | |
| self.word_to_idx = {word: idx for idx, word in enumerate(self.vocab)} | |
| self.data = self._generate_arithmetic_data() | |
| self._save_vocabulary() | |
| def _generate_arithmetic_data(self) -> List[Dict[str, Any]]: | |
| data = [] | |
| for a in range(self.number_range[0], min(self.number_range[1], 21)): | |
| for b in range(self.number_range[0], min(self.number_range[1], 21)): | |
| for op in ["+", "-", "*"]: | |
| if op == "+" and a + b <= 100: | |
| expr = f"{a} {op} {b}" | |
| result = a + b | |
| elif op == "-" and a >= b: | |
| expr = f"{a} {op} {b}" | |
| result = a - b | |
| elif op == "*" and a * b <= 100: | |
| expr = f"{a} {op} {b}" | |
| result = a * b | |
| else: | |
| continue | |
| data.append( | |
| { | |
| "expression": expr, | |
| "result": result, | |
| "expr_tokens": self._tokenize(expr), | |
| "result_tokens": self._tokenize(str(result)), | |
| "complexity": 1, | |
| "is_compositional": False, | |
| } | |
| ) | |
| for a in range(1, 11): | |
| for b in range(1, 11): | |
| for c in range(1, 11): | |
| for op1, op2 in [("*", "+"), ("+", "*"), ("+", "+"), ("*", "*")]: | |
| try: | |
| if op1 == "+" and op2 == "+": | |
| result = (a + b) + c | |
| expr = f"( {a} + {b} ) + {c}" | |
| elif op1 == "*" and op2 == "+": | |
| result = (a * b) + c | |
| expr = f"( {a} * {b} ) + {c}" | |
| elif op1 == "+" and op2 == "*": | |
| result = (a + b) * c | |
| expr = f"( {a} + {b} ) * {c}" | |
| elif op1 == "*" and op2 == "*": | |
| result = (a * b) * c | |
| expr = f"( {a} * {b} ) * {c}" | |
| else: | |
| continue | |
| if result <= 200: | |
| data.append( | |
| { | |
| "expression": expr, | |
| "result": result, | |
| "expr_tokens": self._tokenize(expr), | |
| "result_tokens": self._tokenize(str(result)), | |
| "complexity": 2, | |
| "is_compositional": True, | |
| } | |
| ) | |
| except: | |
| continue | |
| for a in range(1, 6): | |
| for b in range(1, 6): | |
| for c in range(1, 6): | |
| for d in range(1, 6): | |
| try: | |
| result = (a + b) * (c + d) | |
| expr = f"( {a} + {b} ) * ( {c} + {d} )" | |
| if result <= 500: | |
| data.append( | |
| { | |
| "expression": expr, | |
| "result": result, | |
| "expr_tokens": self._tokenize(expr), | |
| "result_tokens": self._tokenize(str(result)), | |
| "complexity": 3, | |
| "is_compositional": True, | |
| } | |
| ) | |
| except: | |
| continue | |
| return data | |
| def _tokenize(self, text: str) -> List[int]: | |
| tokens = text.replace("(", " ( ").replace(")", " ) ").split() | |
| return [self.word_to_idx.get(token, 0) for token in tokens] | |
| def _save_vocabulary(self): | |
| os.makedirs(self.data_dir, exist_ok=True) | |
| vocab_path = os.path.join(self.data_dir, "arithmetic_vocab.json") | |
| vocab_data = { | |
| "vocab": self.vocab, | |
| "word_to_idx": self.word_to_idx, | |
| "vocab_size": self.vocab_size, | |
| "number_range": self.number_range, | |
| } | |
| with open(vocab_path, "w") as f: | |
| json.dump(vocab_data, f, indent=2) | |
| def __len__(self) -> int: | |
| return len(self.data) | |
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: | |
| item = self.data[idx] | |
| max_len = 20 | |
| expr_tokens = item["expr_tokens"][:max_len] | |
| expr_tokens += [0] * (max_len - len(expr_tokens)) | |
| return { | |
| "expression": torch.tensor(expr_tokens, dtype=torch.long), | |
| "result": torch.tensor([item["result"]], dtype=torch.float), | |
| "complexity": item["complexity"], | |
| "is_compositional": item["is_compositional"], | |
| "expr_text": item["expression"], | |
| } | |
| def get_statistics(self) -> Dict[str, Any]: | |
| complexities = [item["complexity"] for item in self.data] | |
| compositional_count = sum(1 for item in self.data if item["is_compositional"]) | |
| return { | |
| "total_examples": len(self.data), | |
| "compositional_examples": compositional_count, | |
| "non_compositional_examples": len(self.data) - compositional_count, | |
| "vocab_size": self.vocab_size, | |
| "avg_expression_length": np.mean( | |
| [len(item["expression"].split()) for item in self.data] | |
| ), | |
| "result_range": ( | |
| min(item["result"] for item in self.data), | |
| max(item["result"] for item in self.data), | |
| ), | |
| "complexity_distribution": { | |
| "1": complexities.count(1), | |
| "2": complexities.count(2), | |
| "3": complexities.count(3), | |
| }, | |
| } | |
| class VisualReasoningDataset(Dataset): | |
| def __init__( | |
| self, num_objects: int = 10, num_properties: int = 5, data_dir: str = None | |
| ): | |
| self.num_objects = num_objects | |
| self.num_properties = num_properties | |
| self.data_dir = data_dir or "data" | |
| self.colors = ["red", "blue", "green", "yellow", "purple"] | |
| self.shapes = ["circle", "square", "triangle", "rectangle", "oval"] | |
| self.sizes = ["small", "medium", "large"] | |
| self.positions = ["left", "right", "above", "below", "center"] | |
| self.vocab = ( | |
| ["<PAD>", "<SOS>", "<EOS>"] | |
| + self.colors | |
| + self.shapes | |
| + self.sizes | |
| + self.positions | |
| + ["of", "and", "is"] | |
| ) | |
| self.vocab_size = len(self.vocab) | |
| self.word_to_idx = {word: idx for idx, word in enumerate(self.vocab)} | |
| self.data = self._generate_visual_data() | |
| self._save_vocabulary() | |
| def _generate_visual_data(self) -> List[Dict[str, Any]]: | |
| data = [] | |
| for color1, shape1, color2, shape2 in itertools.product( | |
| self.colors[:3], self.shapes[:3], self.colors[:3], self.shapes[:3] | |
| ): | |
| for position in self.positions[:4]: | |
| if color1 != color2 or shape1 != shape2: | |
| obj1 = f"{color1} {shape1}" | |
| obj2 = f"{color2} {shape2}" | |
| description = f"{obj1} is {position} of {obj2}" | |
| data.append( | |
| { | |
| "description": description, | |
| "description_tokens": self._tokenize(description), | |
| "object1": obj1, | |
| "object2": obj2, | |
| "relation": position, | |
| "complexity": 1, | |
| "is_compositional": True, | |
| } | |
| ) | |
| for color, shape, size in itertools.product( | |
| self.colors[:3], self.shapes[:3], self.sizes | |
| ): | |
| description = f"{size} {color} {shape}" | |
| data.append( | |
| { | |
| "description": description, | |
| "description_tokens": self._tokenize(description), | |
| "object1": description, | |
| "object2": None, | |
| "relation": None, | |
| "complexity": 2, | |
| "is_compositional": True, | |
| } | |
| ) | |
| for color1, shape1, color2, shape2 in itertools.product( | |
| self.colors[:2], self.shapes[:2], self.colors[:2], self.shapes[:2] | |
| ): | |
| if color1 != color2 or shape1 != shape2: | |
| description = f"{color1} {shape1} and {color2} {shape2}" | |
| data.append( | |
| { | |
| "description": description, | |
| "description_tokens": self._tokenize(description), | |
| "object1": f"{color1} {shape1}", | |
| "object2": f"{color2} {shape2}", | |
| "relation": "and", | |
| "complexity": 3, | |
| "is_compositional": True, | |
| } | |
| ) | |
| return data | |
| def _tokenize(self, text: str) -> List[int]: | |
| words = text.lower().split() | |
| return [self.word_to_idx.get(word, 0) for word in words] | |
| def _save_vocabulary(self): | |
| os.makedirs(self.data_dir, exist_ok=True) | |
| vocab_path = os.path.join(self.data_dir, "visual_vocab.json") | |
| vocab_data = { | |
| "vocab": self.vocab, | |
| "word_to_idx": self.word_to_idx, | |
| "vocab_size": self.vocab_size, | |
| } | |
| with open(vocab_path, "w") as f: | |
| json.dump(vocab_data, f, indent=2) | |
| def __len__(self) -> int: | |
| return len(self.data) | |
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: | |
| item = self.data[idx] | |
| max_len = 15 | |
| desc_tokens = item["description_tokens"][:max_len] | |
| desc_tokens += [0] * (max_len - len(desc_tokens)) | |
| return { | |
| "description": torch.tensor(desc_tokens, dtype=torch.long), | |
| "description_text": item["description"], | |
| "object1": item["object1"], | |
| "object2": item["object2"], | |
| "relation": item["relation"], | |
| "complexity": item["complexity"], | |
| "is_compositional": item["is_compositional"], | |
| } | |
| def get_statistics(self) -> Dict[str, Any]: | |
| complexities = [item["complexity"] for item in self.data] | |
| compositional_count = sum(1 for item in self.data if item["is_compositional"]) | |
| return { | |
| "total_examples": len(self.data), | |
| "compositional_examples": compositional_count, | |
| "non_compositional_examples": len(self.data) - compositional_count, | |
| "vocab_size": self.vocab_size, | |
| "avg_description_length": np.mean( | |
| [len(item["description"].split()) for item in self.data] | |
| ), | |
| "complexity_distribution": { | |
| "1": complexities.count(1), | |
| "2": complexities.count(2), | |
| "3": complexities.count(3), | |
| }, | |
| } | |
| class SystematicSplitGenerator: | |
| def create_compositional_split( | |
| dataset: Dataset, held_out_combinations: List[str], split_ratio: float = 0.8 | |
| ) -> Tuple[List[int], List[int]]: | |
| train_indices = [] | |
| test_indices = [] | |
| for idx in range(len(dataset)): | |
| item = dataset[idx] | |
| item_text = ( | |
| item.get("command_text", "") | |
| or item.get("expr_text", "") | |
| or item.get("description_text", "") | |
| ) | |
| is_held_out = any(combo in item_text for combo in held_out_combinations) | |
| if is_held_out: | |
| test_indices.append(idx) | |
| else: | |
| train_indices.append(idx) | |
| return train_indices, test_indices | |
| def create_length_split( | |
| dataset: Dataset, length_threshold: int, split_ratio: float = 0.8 | |
| ) -> Tuple[List[int], List[int]]: | |
| train_indices = [] | |
| test_indices = [] | |
| for idx in range(len(dataset)): | |
| item = dataset[idx] | |
| if "command" in item: | |
| seq_len = (item["command"] != 0).sum().item() | |
| elif "expression" in item: | |
| seq_len = (item["expression"] != 0).sum().item() | |
| elif "description" in item: | |
| seq_len = (item["description"] != 0).sum().item() | |
| else: | |
| seq_len = 0 | |
| if seq_len <= length_threshold: | |
| train_indices.append(idx) | |
| else: | |
| test_indices.append(idx) | |
| return train_indices, test_indices | |
| def create_complexity_split( | |
| dataset: Dataset, complexity_threshold: int, split_ratio: float = 0.8 | |
| ) -> Tuple[List[int], List[int]]: | |
| train_indices = [] | |
| test_indices = [] | |
| for idx in range(len(dataset)): | |
| item = dataset[idx] | |
| complexity = item.get("complexity", 1) | |
| if complexity <= complexity_threshold: | |
| train_indices.append(idx) | |
| else: | |
| test_indices.append(idx) | |
| return train_indices, test_indices | |
| def create_data_loader( | |
| dataset: Dataset, indices: List[int], batch_size: int = 32, shuffle: bool = True | |
| ) -> DataLoader: | |
| class SubsetDataset(Dataset): | |
| def __init__(self, dataset, indices): | |
| self.dataset = dataset | |
| self.indices = indices | |
| def __len__(self): | |
| return len(self.indices) | |
| def __getitem__(self, idx): | |
| return self.dataset[self.indices[idx]] | |
| subset = SubsetDataset(dataset, indices) | |
| return DataLoader(subset, batch_size=batch_size, shuffle=shuffle) | |
Xet Storage Details
- Size:
- 24 kB
- Xet hash:
- e67687b5617c6d872afc05d05963f980f673b049df3ac92406abcf2c74974b2b
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.