CoolFace
Datasetpublic

Kalso42/WorldModelForMaze

WorldModelForMaze Code, datasets, and trained checkpoints for studying world-model representations in maze navigation, based on a modified NanoGPT. Contents *.py — training, testing, probing, and visualization scripts (see readme.md). model/ — architectures: transformer, transformer-rope, transformer-nextlat, mamba, mamba2, gated-deltanet, gru. data/maze/100/ — tokenized maze datasets for Tasks A/C/E/H/I (RWs paths, 100 nodes). out/ — final (10000-iter)… See the full description on the dataset page: https://huggingface.co/datasets/Kalso42/WorldModelForMaze.

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes627downloads
prepare_minigpt.py165 linesDownload Raw Back to maze
1import os
2import pickle
3import numpy as np
4import re
5import argparse
6
7parser = argparse.ArgumentParser(description='Create the dataset based on the given parameters.')  
8parser.add_argument('--num_nodes', type=int, default=100, help='Number of nodes in the graph')  
9parser.add_argument('--num_of_paths', type=int, default=20, help='Number of paths per pair nodes in training dataset')  
10args = parser.parse_args()  
11
12num_nodes = args.num_nodes
13
14if(args.num_of_paths == 0):
15    train_file_path = os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/train.txt')
16    val_file_path = os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/test.txt')
17else:
18    train_file_path = os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/train_{args.num_of_paths}.txt')
19    val_file_path = os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/test.txt')
20# test_file_path = os.path.join(os.path.dirname(__file__), 'test.txt')
21
22with open(train_file_path, 'r') as f:
23    train_data = f.read()
24print(f"length of train dataset in characters: {len(train_data):,}")
25
26with open(val_file_path, 'r') as f:
27    val_data = f.read()
28print(f"length of val dataset in characters: {len(val_data):,}")
29
30all_data = train_data + val_data
31
32def find_characters(data_string):
33    pattern = r'\d+|\D'
34    matches = re.findall(pattern, data_string)
35    return set(matches)
36
37def process_reasoning(s):
38    split_text = s.split('\n')
39    #split_text = [s + '\n' for s in split_text if s != ""]
40    ret = []
41    for st in split_text:
42        if(st != ""):
43            enc_str = encode(st) + [1]
44            ret += enc_str +[0] * (block_size + 1 - len(enc_str))
45    return ret
46
47def get_block_size(s):
48    split_text = s.split('\n')
49    #split_text = [s + '\n' for s in split_text if s != ""]
50    ret = []
51    bs = 0
52    for st in split_text:
53        if(st != ""):
54            enc_str = encode(st) + [1]
55            bs = max(bs, len(enc_str))
56    return bs
57
58
59def encode_string(s, stonum):
60    ss = s.split(" ")
61    encoded_string = [stonum[ch] for ch in ss]
62    return encoded_string
63
64def decode_string(l, numtos):
65    dec = ""
66    for i in l:
67        dec = dec + numtos[i] + " "
68    return dec[:-1]
69
70
71# get all the unique characters that occur in this text
72chars = sorted(list(find_characters(all_data)))
73# direction tokens for maze paths
74direction_tokens = ['N','S','E','W']
75# task tokens for multi-task support
76task_tokens = ['A', 'B', 'C', 'D', 'E', 'F', 'G']
77# special tokens: 'x' marks a wall-hit / unreachable (wrong-path) terminator
78special_tokens = ['x']
79# vocab = node ids + PAD + newline + direction tokens + task tokens + special tokens
80vocab_size = num_nodes + 2 + len(direction_tokens) + len(task_tokens) + len(special_tokens)
81print("all the unique characters:", ' '.join(chars))
82print(f"vocab size: {vocab_size:,}")
83
84# create a mapping from characters to integers
85stoi = {}
86itos = {}
87
88for i in range(num_nodes):
89    stoi[str(i)] = i+2
90    itos[i+2] = str(i)
91
92# map direction tokens after the node id tokens
93base = 2 + num_nodes
94for idx, tok in enumerate(direction_tokens):
95    stoi[tok] = base + idx
96    itos[base + idx] = tok
97
98# map task tokens after direction tokens
99base = 2 + num_nodes + len(direction_tokens)
100for idx, tok in enumerate(task_tokens):
101    stoi[tok] = base + idx
102    itos[base + idx] = tok
103
104# map special tokens (e.g. 'x') after task tokens
105base = 2 + num_nodes + len(direction_tokens) + len(task_tokens)
106for idx, tok in enumerate(special_tokens):
107    stoi[tok] = base + idx
108    itos[base + idx] = tok
109
110stoi['[PAD]'] = 0
111itos[0] = '[PAD]'
112stoi['\n'] = 1
113itos[1] = '\n'
114
115def encode(s):
116    return encode_string(s, stoi) # encoder: take a string, output a list of integers
117def decode(l):
118    return decode_string(l, itos) # decoder: take a list of integers, output a string
119
120# encode both to integers
121block_size = (max(get_block_size(train_data), get_block_size(val_data)) // 32 + 1) * 32
122
123print(f"the block size is {block_size}")
124
125train_ids = process_reasoning(train_data)
126
127val_ids = process_reasoning(val_data)
128
129print(f"train has {len(train_ids):,} tokens")
130print(f"val has {len(val_ids):,} tokens")
131
132# export to bin files
133train_ids = np.array(train_ids, dtype=np.uint16)
134val_ids = np.array(val_ids, dtype=np.uint16)
135
136if(args.num_of_paths == 0):
137    train_ids.tofile(os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/train.bin'))
138    val_ids.tofile(os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/val.bin'))
139else:
140    train_ids.tofile(os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/train_{args.num_of_paths}.bin'))
141    val_ids.tofile(os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/val.bin'))
142
143
144unreachable = False; simple_format = True
145if 'x' in chars:
146    unreachable = True
147if ':' in chars:
148    simple_format = False
149    
150
151# save the meta information as well, to help us encode/decode later
152meta = {
153    'unreachable': unreachable,
154    'simple_format': simple_format,
155    'block_size': block_size,
156    'vocab_size': vocab_size,
157    'itos': itos,
158    'stoi': stoi,
159}
160
161print(stoi)
162print(itos)
163with open(os.path.join(os.path.dirname(__file__), f'{args.num_nodes}/meta.pkl'), 'wb') as f:
164    pickle.dump(meta, f)
165