tahamajs's picture
download
raw
24 kB
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:
@staticmethod
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
@staticmethod
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
@staticmethod
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.