Buckets:
tahamajs/Sysmem2_in_AI / ComputerAssignments /CA6_systematic_generalization /src /datasets /advanced_datasets.py
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import Dataset, DataLoader | |
| import numpy as np | |
| import json | |
| import random | |
| from typing import Dict, List, Tuple, Optional, Any, Union | |
| from dataclasses import dataclass | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| from abc import ABC, abstractmethod | |
| import itertools | |
| from collections import defaultdict | |
| import re | |
| class DatasetConfig: | |
| name: str | |
| input_dim: int | |
| output_dim: int | |
| num_examples: int | |
| composition_depth: int | |
| complexity_levels: List[int] | |
| domain_types: List[str] | |
| adversarial_ratio: float = 0.1 | |
| noise_level: float = 0.05 | |
| class AdvancedCompositionalDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, split: str = 'train'): | |
| self.config = config | |
| self.split = split | |
| self.examples = [] | |
| self.composition_rules = self._initialize_composition_rules() | |
| self.domain_vocabularies = self._initialize_domain_vocabularies() | |
| self._generate_examples() | |
| def _initialize_composition_rules(self) -> Dict[str, List[str]]: | |
| return { | |
| 'arithmetic': [ | |
| 'add', 'subtract', 'multiply', 'divide', | |
| 'power', 'sqrt', 'log', 'exp' | |
| ], | |
| 'logical': [ | |
| 'and', 'or', 'not', 'xor', 'implies', 'iff' | |
| ], | |
| 'spatial': [ | |
| 'left_of', 'right_of', 'above', 'below', | |
| 'inside', 'outside', 'near', 'far' | |
| ], | |
| 'temporal': [ | |
| 'before', 'after', 'during', 'simultaneous', | |
| 'precedes', 'follows', 'overlaps' | |
| ], | |
| 'causal': [ | |
| 'causes', 'prevents', 'enables', 'inhibits', | |
| 'triggers', 'blocks', 'facilitates' | |
| ] | |
| } | |
| def _initialize_domain_vocabularies(self) -> Dict[str, List[str]]: | |
| return { | |
| 'arithmetic': [str(i) for i in range(100)], | |
| 'logical': ['true', 'false', 'A', 'B', 'C', 'D'], | |
| 'spatial': ['circle', 'square', 'triangle', 'rectangle', 'red', 'blue', 'green'], | |
| 'temporal': ['morning', 'afternoon', 'evening', 'night', 'day', 'week', 'month'], | |
| 'causal': ['rain', 'sun', 'wind', 'storm', 'growth', 'decay', 'change'] | |
| } | |
| def _generate_examples(self): | |
| num_examples = self.config.num_examples | |
| for i in range(num_examples): | |
| domain = random.choice(self.config.domain_types) | |
| complexity = random.choice(self.config.complexity_levels) | |
| if domain == 'arithmetic': | |
| example = self._generate_arithmetic_example(complexity) | |
| elif domain == 'logical': | |
| example = self._generate_logical_example(complexity) | |
| elif domain == 'spatial': | |
| example = self._generate_spatial_example(complexity) | |
| elif domain == 'temporal': | |
| example = self._generate_temporal_example(complexity) | |
| elif domain == 'causal': | |
| example = self._generate_causal_example(complexity) | |
| else: | |
| example = self._generate_mixed_example(complexity) | |
| if self.config.noise_level > 0: | |
| example = self._add_noise(example) | |
| self.examples.append(example) | |
| def _generate_arithmetic_example(self, complexity: int) -> Dict[str, Any]: | |
| numbers = random.sample(self.domain_vocabularies['arithmetic'], complexity + 2) | |
| operations = random.sample(self.composition_rules['arithmetic'], complexity) | |
| expression_parts = [numbers[0]] | |
| for i in range(complexity): | |
| expression_parts.append(operations[i]) | |
| expression_parts.append(numbers[i + 1]) | |
| expression = ' '.join(expression_parts) | |
| try: | |
| result = eval(expression.replace(' ', '')) | |
| except: | |
| result = 0 | |
| return { | |
| 'input': expression, | |
| 'output': str(result), | |
| 'domain': 'arithmetic', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(expression) | |
| } | |
| def _generate_logical_example(self, complexity: int) -> Dict[str, Any]: | |
| variables = random.sample(self.domain_vocabularies['logical'], min(complexity + 1, 4)) | |
| operations = random.sample(self.composition_rules['logical'], complexity) | |
| expression_parts = [variables[0]] | |
| for i in range(complexity): | |
| expression_parts.append(operations[i]) | |
| expression_parts.append(variables[min(i + 1, len(variables) - 1)]) | |
| expression = ' '.join(expression_parts) | |
| result = self._evaluate_logical_expression(expression) | |
| return { | |
| 'input': expression, | |
| 'output': result, | |
| 'domain': 'logical', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(expression) | |
| } | |
| def _generate_spatial_example(self, complexity: int) -> Dict[str, Any]: | |
| objects = random.sample(self.domain_vocabularies['spatial'], complexity + 1) | |
| relations = random.sample(self.composition_rules['spatial'], complexity) | |
| description_parts = [objects[0]] | |
| for i in range(complexity): | |
| description_parts.append(relations[i]) | |
| description_parts.append(objects[i + 1]) | |
| description = ' '.join(description_parts) | |
| spatial_repr = self._generate_spatial_representation(objects, relations) | |
| return { | |
| 'input': description, | |
| 'output': spatial_repr, | |
| 'domain': 'spatial', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(description) | |
| } | |
| def _generate_temporal_example(self, complexity: int) -> Dict[str, Any]: | |
| events = random.sample(self.domain_vocabularies['temporal'], complexity + 1) | |
| relations = random.sample(self.composition_rules['temporal'], complexity) | |
| sequence_parts = [events[0]] | |
| for i in range(complexity): | |
| sequence_parts.append(relations[i]) | |
| sequence_parts.append(events[i + 1]) | |
| sequence = ' '.join(sequence_parts) | |
| temporal_repr = self._generate_temporal_representation(events, relations) | |
| return { | |
| 'input': sequence, | |
| 'output': temporal_repr, | |
| 'domain': 'temporal', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(sequence) | |
| } | |
| def _generate_causal_example(self, complexity: int) -> Dict[str, Any]: | |
| causes = random.sample(self.domain_vocabularies['causal'], complexity + 1) | |
| relations = random.sample(self.composition_rules['causal'], complexity) | |
| chain_parts = [causes[0]] | |
| for i in range(complexity): | |
| chain_parts.append(relations[i]) | |
| chain_parts.append(causes[i + 1]) | |
| chain = ' '.join(chain_parts) | |
| causal_repr = self._generate_causal_representation(causes, relations) | |
| return { | |
| 'input': chain, | |
| 'output': causal_repr, | |
| 'domain': 'causal', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(chain) | |
| } | |
| def _generate_mixed_example(self, complexity: int) -> Dict[str, Any]: | |
| domains = random.sample(self.config.domain_types, min(2, len(self.config.domain_types))) | |
| mixed_parts = [] | |
| for domain in domains: | |
| if domain == 'arithmetic': | |
| part = self._generate_arithmetic_example(1) | |
| elif domain == 'logical': | |
| part = self._generate_logical_example(1) | |
| elif domain == 'spatial': | |
| part = self._generate_spatial_example(1) | |
| elif domain == 'temporal': | |
| part = self._generate_temporal_example(1) | |
| elif domain == 'causal': | |
| part = self._generate_causal_example(1) | |
| else: | |
| continue | |
| mixed_parts.append(part['input']) | |
| mixed_input = ' AND '.join(mixed_parts) | |
| mixed_output = ' AND '.join([part['output'] for part in mixed_parts]) | |
| return { | |
| 'input': mixed_input, | |
| 'output': mixed_output, | |
| 'domain': 'mixed', | |
| 'complexity': complexity, | |
| 'composition_structure': self._analyze_composition_structure(mixed_input) | |
| } | |
| def _evaluate_logical_expression(self, expression: str) -> str: | |
| if 'true' in expression and 'false' in expression: | |
| return 'false' | |
| elif 'true' in expression: | |
| return 'true' | |
| elif 'false' in expression: | |
| return 'false' | |
| else: | |
| return 'unknown' | |
| def _generate_spatial_representation(self, objects: List[str], relations: List[str]) -> str: | |
| layout = [] | |
| for i, obj in enumerate(objects): | |
| if i < len(relations): | |
| layout.append(f"{obj} {relations[i]}") | |
| else: | |
| layout.append(obj) | |
| return ' | '.join(layout) | |
| def _generate_temporal_representation(self, events: List[str], relations: List[str]) -> str: | |
| sequence = [] | |
| for i, event in enumerate(events): | |
| if i < len(relations): | |
| sequence.append(f"{event} {relations[i]}") | |
| else: | |
| sequence.append(event) | |
| return ' -> '.join(sequence) | |
| def _generate_causal_representation(self, causes: List[str], relations: List[str]) -> str: | |
| chain = [] | |
| for i, cause in enumerate(causes): | |
| if i < len(relations): | |
| chain.append(f"{cause} {relations[i]}") | |
| else: | |
| chain.append(cause) | |
| return ' => '.join(chain) | |
| def _analyze_composition_structure(self, expression: str) -> Dict[str, Any]: | |
| parts = expression.split() | |
| return { | |
| 'length': len(parts), | |
| 'operators': [part for part in parts if part in | |
| sum(self.composition_rules.values(), [])], | |
| 'operands': [part for part in parts if part not in | |
| sum(self.composition_rules.values(), [])], | |
| 'depth': self._calculate_composition_depth(parts) | |
| } | |
| def _calculate_composition_depth(self, parts: List[str]) -> int: | |
| depth = 0 | |
| max_depth = 0 | |
| for part in parts: | |
| if part in sum(self.composition_rules.values(), []): | |
| depth += 1 | |
| max_depth = max(max_depth, depth) | |
| else: | |
| depth = max(0, depth - 1) | |
| return max_depth | |
| def _add_noise(self, example: Dict[str, Any]) -> Dict[str, Any]: | |
| if random.random() < self.config.noise_level: | |
| example['input'] += random.choice(['!', '?', '*']) | |
| return example | |
| def __len__(self): | |
| return len(self.examples) | |
| def __getitem__(self, idx): | |
| example = self.examples[idx] | |
| input_tensor = self._text_to_tensor(example['input']) | |
| output_tensor = self._text_to_tensor(example['output']) | |
| return { | |
| 'input': input_tensor, | |
| 'target': output_tensor, | |
| 'metadata': example | |
| } | |
| def _text_to_tensor(self, text: str) -> torch.Tensor: | |
| chars = list(text) | |
| char_to_idx = {char: idx for idx, char in enumerate(set(chars))} | |
| tensor = torch.zeros(len(chars), len(char_to_idx)) | |
| for i, char in enumerate(chars): | |
| tensor[i, char_to_idx[char]] = 1.0 | |
| return tensor | |
| class HierarchicalReasoningDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, split: str = 'train'): | |
| self.config = config | |
| self.split = split | |
| self.hierarchy_levels = 5 | |
| self.examples = [] | |
| self._generate_hierarchical_examples() | |
| def _generate_hierarchical_examples(self): | |
| for i in range(self.config.num_examples): | |
| hierarchy = self._generate_hierarchy() | |
| task = self._create_reasoning_task(hierarchy) | |
| self.examples.append(task) | |
| def _generate_hierarchy(self) -> Dict[str, Any]: | |
| hierarchy = { | |
| 'level_0': ['A', 'B', 'C'], | |
| 'level_1': ['AB', 'BC', 'AC'], | |
| 'level_2': ['ABC', 'BCA', 'CAB'], | |
| 'level_3': ['ABCA', 'BCAB', 'CABC'], | |
| 'level_4': ['ABCAB', 'BCABC', 'CABCA'] | |
| } | |
| return hierarchy | |
| def _create_reasoning_task(self, hierarchy: Dict[str, Any]) -> Dict[str, Any]: | |
| level = random.choice(list(hierarchy.keys())) | |
| elements = hierarchy[level] | |
| question = f"What is the composition of {random.choice(elements)}?" | |
| answer = self._generate_hierarchical_answer(elements, level) | |
| return { | |
| 'input': question, | |
| 'output': answer, | |
| 'hierarchy_level': level, | |
| 'elements': elements, | |
| 'reasoning_steps': self._generate_reasoning_steps(elements, level) | |
| } | |
| def _generate_hierarchical_answer(self, elements: List[str], level: str) -> str: | |
| if level == 'level_0': | |
| return f"Base elements: {', '.join(elements)}" | |
| elif level == 'level_1': | |
| return f"First-level composition: {', '.join(elements)}" | |
| elif level == 'level_2': | |
| return f"Second-level composition: {', '.join(elements)}" | |
| elif level == 'level_3': | |
| return f"Third-level composition: {', '.join(elements)}" | |
| else: | |
| return f"Fourth-level composition: {', '.join(elements)}" | |
| def _generate_reasoning_steps(self, elements: List[str], level: str) -> List[str]: | |
| steps = [] | |
| if level != 'level_0': | |
| steps.append(f"Identify base elements in {level}") | |
| steps.append(f"Apply composition rules") | |
| steps.append(f"Generate composite elements") | |
| return steps | |
| def __len__(self): | |
| return len(self.examples) | |
| def __getitem__(self, idx): | |
| example = self.examples[idx] | |
| return { | |
| 'input': example['input'], | |
| 'target': example['output'], | |
| 'metadata': example | |
| } | |
| class CrossDomainTransferDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, source_domain: str, target_domain: str): | |
| self.config = config | |
| self.source_domain = source_domain | |
| self.target_domain = target_domain | |
| self.examples = [] | |
| self._generate_cross_domain_examples() | |
| def _generate_cross_domain_examples(self): | |
| for i in range(self.config.num_examples): | |
| source_example = self._generate_domain_example(self.source_domain) | |
| target_example = self._generate_domain_example(self.target_domain) | |
| transfer_mapping = self._create_transfer_mapping(source_example, target_example) | |
| self.examples.append({ | |
| 'source': source_example, | |
| 'target': target_example, | |
| 'transfer_mapping': transfer_mapping, | |
| 'transfer_difficulty': self._calculate_transfer_difficulty(source_example, target_example) | |
| }) | |
| def _generate_domain_example(self, domain: str) -> Dict[str, Any]: | |
| if domain == 'arithmetic': | |
| return self._generate_arithmetic_example() | |
| elif domain == 'logical': | |
| return self._generate_logical_example() | |
| elif domain == 'spatial': | |
| return self._generate_spatial_example() | |
| else: | |
| return self._generate_generic_example() | |
| def _generate_arithmetic_example(self) -> Dict[str, Any]: | |
| a, b = random.randint(1, 10), random.randint(1, 10) | |
| operation = random.choice(['+', '-', '*', '/']) | |
| if operation == '+': | |
| result = a + b | |
| elif operation == '-': | |
| result = a - b | |
| elif operation == '*': | |
| result = a * b | |
| else: | |
| result = a / b if b != 0 else 0 | |
| return { | |
| 'input': f"{a} {operation} {b}", | |
| 'output': str(result), | |
| 'domain': 'arithmetic' | |
| } | |
| def _generate_logical_example(self) -> Dict[str, Any]: | |
| p, q = random.choice([True, False]), random.choice([True, False]) | |
| operation = random.choice(['and', 'or', 'xor']) | |
| if operation == 'and': | |
| result = p and q | |
| elif operation == 'or': | |
| result = p or q | |
| else: | |
| result = p != q | |
| return { | |
| 'input': f"{p} {operation} {q}", | |
| 'output': str(result), | |
| 'domain': 'logical' | |
| } | |
| def _generate_spatial_example(self) -> Dict[str, Any]: | |
| objects = ['circle', 'square', 'triangle'] | |
| relations = ['left_of', 'right_of', 'above', 'below'] | |
| obj1, obj2 = random.sample(objects, 2) | |
| relation = random.choice(relations) | |
| return { | |
| 'input': f"{obj1} {relation} {obj2}", | |
| 'output': f"spatial_layout_{relation}", | |
| 'domain': 'spatial' | |
| } | |
| def _generate_generic_example(self) -> Dict[str, Any]: | |
| return { | |
| 'input': f"generic_input_{random.randint(1, 100)}", | |
| 'output': f"generic_output_{random.randint(1, 100)}", | |
| 'domain': 'generic' | |
| } | |
| def _create_transfer_mapping(self, source: Dict[str, Any], target: Dict[str, Any]) -> Dict[str, str]: | |
| return { | |
| 'input_mapping': f"{source['input']} -> {target['input']}", | |
| 'output_mapping': f"{source['output']} -> {target['output']}", | |
| 'domain_mapping': f"{source['domain']} -> {target['domain']}" | |
| } | |
| def _calculate_transfer_difficulty(self, source: Dict[str, Any], target: Dict[str, Any]) -> float: | |
| domain_similarity = { | |
| ('arithmetic', 'logical'): 0.3, | |
| ('arithmetic', 'spatial'): 0.1, | |
| ('logical', 'spatial'): 0.2, | |
| ('arithmetic', 'arithmetic'): 1.0, | |
| ('logical', 'logical'): 1.0, | |
| ('spatial', 'spatial'): 1.0 | |
| } | |
| key = (source['domain'], target['domain']) | |
| return domain_similarity.get(key, 0.5) | |
| def __len__(self): | |
| return len(self.examples) | |
| def __getitem__(self, idx): | |
| example = self.examples[idx] | |
| return { | |
| 'source_input': example['source']['input'], | |
| 'source_target': example['source']['output'], | |
| 'target_input': example['target']['input'], | |
| 'target_target': example['target']['output'], | |
| 'metadata': example | |
| } | |
| class AdversarialCompositionalDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, adversarial_ratio: float = 0.3): | |
| self.config = config | |
| self.adversarial_ratio = adversarial_ratio | |
| self.examples = [] | |
| self._generate_adversarial_examples() | |
| def _generate_adversarial_examples(self): | |
| for i in range(self.config.num_examples): | |
| if random.random() < self.adversarial_ratio: | |
| example = self._generate_adversarial_example() | |
| else: | |
| example = self._generate_normal_example() | |
| self.examples.append(example) | |
| def _generate_adversarial_example(self) -> Dict[str, Any]: | |
| misleading_input = self._create_misleading_input() | |
| correct_output = self._generate_correct_output(misleading_input) | |
| misleading_output = self._generate_misleading_output(misleading_input) | |
| return { | |
| 'input': misleading_input, | |
| 'correct_output': correct_output, | |
| 'misleading_output': misleading_output, | |
| 'is_adversarial': True, | |
| 'adversarial_type': self._classify_adversarial_type(misleading_input) | |
| } | |
| def _generate_normal_example(self) -> Dict[str, Any]: | |
| normal_input = self._create_normal_input() | |
| normal_output = self._generate_correct_output(normal_input) | |
| return { | |
| 'input': normal_input, | |
| 'correct_output': normal_output, | |
| 'misleading_output': None, | |
| 'is_adversarial': False, | |
| 'adversarial_type': None | |
| } | |
| def _create_misleading_input(self) -> str: | |
| misleading_patterns = [ | |
| "2 + 2 = 5", | |
| "true and false = true", | |
| "circle left_of square = right_of", | |
| "morning after evening = before" | |
| ] | |
| return random.choice(misleading_patterns) | |
| def _create_normal_input(self) -> str: | |
| normal_patterns = [ | |
| "2 + 2 = 4", | |
| "true and false = false", | |
| "circle left_of square = left_of", | |
| "morning after evening = after" | |
| ] | |
| return random.choice(normal_patterns) | |
| def _generate_correct_output(self, input_text: str) -> str: | |
| if "2 + 2 = 4" in input_text: | |
| return "4" | |
| elif "true and false = false" in input_text: | |
| return "false" | |
| elif "left_of" in input_text: | |
| return "left_of" | |
| elif "after" in input_text: | |
| return "after" | |
| else: | |
| return "correct" | |
| def _generate_misleading_output(self, input_text: str) -> str: | |
| if "2 + 2 = 5" in input_text: | |
| return "5" | |
| elif "true and false = true" in input_text: | |
| return "true" | |
| elif "right_of" in input_text: | |
| return "right_of" | |
| elif "before" in input_text: | |
| return "before" | |
| else: | |
| return "misleading" | |
| def _classify_adversarial_type(self, input_text: str) -> str: | |
| if any(op in input_text for op in ['+', '-', '*', '/']): | |
| return 'arithmetic_adversarial' | |
| elif any(op in input_text for op in ['and', 'or', 'not']): | |
| return 'logical_adversarial' | |
| elif any(op in input_text for op in ['left_of', 'right_of', 'above', 'below']): | |
| return 'spatial_adversarial' | |
| else: | |
| return 'generic_adversarial' | |
| def __len__(self): | |
| return len(self.examples) | |
| def __getitem__(self, idx): | |
| example = self.examples[idx] | |
| return { | |
| 'input': example['input'], | |
| 'target': example['correct_output'], | |
| 'adversarial_target': example['misleading_output'], | |
| 'is_adversarial': example['is_adversarial'], | |
| 'metadata': example | |
| } | |
| class DynamicCompositionalDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, dynamic_level: float = 0.5): | |
| self.config = config | |
| self.dynamic_level = dynamic_level | |
| self.examples = [] | |
| self.composition_rules = self._initialize_dynamic_rules() | |
| self._generate_dynamic_examples() | |
| def _initialize_dynamic_rules(self) -> Dict[str, List[str]]: | |
| return { | |
| 'adaptive': ['adapt', 'adjust', 'modify', 'change'], | |
| 'temporal': ['evolve', 'develop', 'progress', 'advance'], | |
| 'contextual': ['context', 'situation', 'environment', 'setting'], | |
| 'interactive': ['interact', 'respond', 'react', 'engage'] | |
| } | |
| def _generate_dynamic_examples(self): | |
| for i in range(self.config.num_examples): | |
| dynamic_example = self._generate_dynamic_example() | |
| self.examples.append(dynamic_example) | |
| def _generate_dynamic_example(self) -> Dict[str, Any]: | |
| rule_type = random.choice(list(self.composition_rules.keys())) | |
| rules = self.composition_rules[rule_type] | |
| base_elements = ['A', 'B', 'C', 'D'] | |
| dynamic_rule = random.choice(rules) | |
| dynamic_input = f"{random.choice(base_elements)} {dynamic_rule} {random.choice(base_elements)}" | |
| dynamic_output = self._generate_dynamic_output(dynamic_input, rule_type) | |
| return { | |
| 'input': dynamic_input, | |
| 'output': dynamic_output, | |
| 'rule_type': rule_type, | |
| 'dynamic_level': self.dynamic_level, | |
| 'adaptation_steps': self._generate_adaptation_steps(dynamic_input, rule_type) | |
| } | |
| def _generate_dynamic_output(self, input_text: str, rule_type: str) -> str: | |
| if rule_type == 'adaptive': | |
| return f"adapted_{input_text.replace(' ', '_')}" | |
| elif rule_type == 'temporal': | |
| return f"evolved_{input_text.replace(' ', '_')}" | |
| elif rule_type == 'contextual': | |
| return f"contextualized_{input_text.replace(' ', '_')}" | |
| elif rule_type == 'interactive': | |
| return f"interactive_{input_text.replace(' ', '_')}" | |
| else: | |
| return f"dynamic_{input_text.replace(' ', '_')}" | |
| def _generate_adaptation_steps(self, input_text: str, rule_type: str) -> List[str]: | |
| steps = [] | |
| if rule_type == 'adaptive': | |
| steps = ["Analyze input", "Identify adaptation needs", "Apply adaptation", "Verify result"] | |
| elif rule_type == 'temporal': | |
| steps = ["Initialize state", "Apply temporal rule", "Update state", "Generate output"] | |
| elif rule_type == 'contextual': | |
| steps = ["Extract context", "Apply contextual rule", "Integrate context", "Generate output"] | |
| elif rule_type == 'interactive': | |
| steps = ["Receive input", "Process interaction", "Generate response", "Update state"] | |
| return steps | |
| def __len__(self): | |
| return len(self.examples) | |
| def __getitem__(self, idx): | |
| example = self.examples[idx] | |
| return { | |
| 'input': example['input'], | |
| 'target': example['output'], | |
| 'metadata': example | |
| } | |
| class MetaLearningDataset(Dataset): | |
| def __init__(self, config: DatasetConfig, num_tasks: int = 10): | |
| self.config = config | |
| self.num_tasks = num_tasks | |
| self.tasks = [] | |
| self._generate_meta_learning_tasks() | |
| def _generate_meta_learning_tasks(self): | |
| for i in range(self.num_tasks): | |
| task = self._generate_meta_task() | |
| self.tasks.append(task) | |
| def _generate_meta_task(self) -> Dict[str, Any]: | |
| support_set = self._generate_support_set() | |
| query_set = self._generate_query_set() | |
| task_description = self._generate_task_description() | |
| return { | |
| 'task_id': len(self.tasks), | |
| 'description': task_description, | |
| 'support_set': support_set, | |
| 'query_set': query_set, | |
| 'task_type': self._classify_task_type(support_set, query_set) | |
| } | |
| def _generate_support_set(self) -> List[Dict[str, Any]]: | |
| support_examples = [] | |
| for i in range(5): | |
| example = { | |
| 'input': f"support_input_{i}", | |
| 'output': f"support_output_{i}", | |
| 'example_id': i | |
| } | |
| support_examples.append(example) | |
| return support_examples | |
| def _generate_query_set(self) -> List[Dict[str, Any]]: | |
| query_examples = [] | |
| for i in range(10): | |
| example = { | |
| 'input': f"query_input_{i}", | |
| 'output': f"query_output_{i}", | |
| 'example_id': i | |
| } | |
| query_examples.append(example) | |
| return query_examples | |
| def _generate_task_description(self) -> str: | |
| descriptions = [ | |
| "Learn to compose arithmetic expressions", | |
| "Learn to compose logical expressions", | |
| "Learn to compose spatial relations", | |
| "Learn to compose temporal sequences", | |
| "Learn to compose causal chains" | |
| ] | |
| return random.choice(descriptions) | |
| def _classify_task_type(self, support_set: List[Dict[str, Any]], | |
| query_set: List[Dict[str, Any]]) -> str: | |
| if 'arithmetic' in support_set[0]['input']: | |
| return 'arithmetic_composition' | |
| elif 'logical' in support_set[0]['input']: | |
| return 'logical_composition' | |
| elif 'spatial' in support_set[0]['input']: | |
| return 'spatial_composition' | |
| else: | |
| return 'generic_composition' | |
| def __len__(self): | |
| return len(self.tasks) | |
| def __getitem__(self, idx): | |
| task = self.tasks[idx] | |
| return { | |
| 'task_id': task['task_id'], | |
| 'description': task['description'], | |
| 'support_set': task['support_set'], | |
| 'query_set': task['query_set'], | |
| 'task_type': task['task_type'] | |
| } | |
Xet Storage Details
- Size:
- 29.5 kB
- Xet hash:
- 84a234ec675b560be76d61b284887789e4841db0a76a4520d6ba81f953a62838
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.