Views
No views yet
def encode(self, text):
byte_seq = text.encode('utf-8')
return [self.byte_to_index[byte] for byte in byte_seq]
def decode(self, indices):
byte_seq = bytes(self.index_to_byte[index] for index in indices)
return byte_seq.decode('utf-8', errors='ignore')def __call__(self, q, k, position_ids):
cos, sin = self._get_cos_sin(position_ids)
q = (q * cos) + (self._rotate_half(q) * sin)
k = (k * cos) + (self._rotate_half(k) * sin)
return q, k
def _get_cos_sin(self, position_ids):
su_factor = self._short_factor
position_ids_expanded = position_ids[:, None, :]
inv_freq = 1.0 / (su_factor * self.rope_theta**(mx.arange(0, self.dim, 2, dtype=mx.float32) / self.dim))
inv_freq_expanded = mx.repeat(inv_freq[None, :, None], position_ids.shape[0], axis=0)
freqs = (inv_freq_expanded @ position_ids_expanded).transpose(0, 2, 1)
emb = mx.concatenate([freqs, freqs], axis=-1)
cos = mx.expand_dims(mx.cos(emb) * self.scaling_factor, axis=1)
sin = mx.expand_dims(mx.sin(emb) * self.scaling_factor, axis=1)
return cos, sin
def _rotate_half(self, x):
midpoint = x.shape[-1] // 2
x1, x2 = x[..., :midpoint], x[..., midpoint:]
return mx.concatenate([-x2, x1], axis=-1)def __call__(self, x, position_ids, attention_mask, cache, use_recurrent_mode):
B, L, _ = x.shape
qkv = self.qkv_proj(x)
q, k, v = mx.split(qkv, self.chop, axis=-1)
q = q.reshape(B, L, self.n_heads, -1).transpose(0, 2, 1, 3)
k = k.reshape(B, L, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
v = v.reshape(B, L, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
if cache is None:
position_ids = mx.arange(q.shape[2], dtype=mx.float32)[None] if position_ids is None else position_ids
q, k = self.rope(q,k,position_ids)
mask = mx.triu(mx.full((v.shape[2], v.shape[2]), -mx.inf), k=1)
if attention_mask is not None:
mask += mx.where(attention_mask[:, :, None]*attention_mask[:, None, :]==1, 0, -mx.inf)
mask = mx.expand_dims(mask, 1)
else:
mask = mask[None, None]
else:
past_k, past_v, past_p, past_m = cache
position_ids = past_p[:,-1:]+1
mask = mx.pad(past_m[:,:,-1:,:], ((0,0),(0,0),(0,0),(0,1)))
q, k = self.rope(q, k, position_ids)
k = mx.concatenate([past_k, k], axis=2)
v = mx.concatenate([past_v, v], axis=2)
cache = (k, v, position_ids, mask)
w = (q * self.scale) @ k.transpose(0, 1, 3, 2)
w += mask
w = mx.softmax(w, axis=-1)
o = w @ v
o = o.transpose(0, 2, 1, 3).reshape(B, L, -1)
return self.o_proj(o).astype(x.dtype), cachedef __call__(self, x, position_ids, attention_mask, cache, use_recurrent_mode):
if use_recurrent_mode:
return self.recurrent_mode(x, cache)
B, L, _ = x.shape
qkv = self.qkv_proj(x)
q, k, v = mx.split(qkv, self.chop, axis=-1)
q = q.reshape(B, L, self.n_heads, -1).transpose(0, 2, 1, 3)
k = k.reshape(B, L, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
v = v.reshape(B, L, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
position_ids = mx.arange(q.shape[2], dtype=mx.float32)[None] if position_ids is None else position_ids
q, k = self.rope(q,k,position_ids)
cache = None
w = (q * self.scale) @ k.transpose(0, 1, 3, 2)
w = w * self._decay(L)
o = w @ v
o = o.transpose(0, 2, 1, 3).reshape(B*L, -1)
o = self.gn(o).reshape(B, L, -1)
return self.o_proj(o).astype(x.dtype), cache
def recurrent_mode(self, x, cache):
if cache is None:
s = mx.zeros((1, 32, 96, 96))
n = 0
else:
s, n = cache
qkv = self.qkv_proj(x)
q, k, v = mx.split(qkv, self.chop, axis=-1)
q = q.reshape(1, 1, self.n_heads, -1).transpose(0, 2, 1, 3)
k = k.reshape(1, 1, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
v = v.reshape(1, 1, self.n_kv_heads, -1).transpose(0, 2, 1, 3)
position_ids = mx.array([[n]])
q, k = self.rope(q,k,position_ids)
k = k * self.scale
s = self._gamma[None, :, None, None] * s + (k.transpose(0, 1, 3, 2) @ v)
o = q @ s
o = o.transpose(0, 2, 1, 3).reshape(1, -1)
o = self.gn(o).reshape(1, 1, -1)
o = self.o_proj(o).astype(x.dtype)
return o, (s, n+1)
def _decay(self, sequence_length):
n = mx.arange(sequence_length)[:,None]
m = mx.arange(sequence_length)[None]
D = (self._gamma[:, None, None] ** (n-m)) * (n >= m)
return Ddef __call__(self, x):
x = self.gate_up_proj(x)
gate, x = mx.split(x, 2, axis=-1)
return self.down_proj(nn.silu(gate) * x)def __call__(self, x, position_ids, attention_mask, cache, use_recurrent_mode):
r, cache = self.self_attn(self.input_layernorm(x), position_ids, attention_mask, cache, use_recurrent_mode)
h = x + r
r = self.mlp(self.post_attention_layernorm(h))
return h + r, cachedef __call__(self, input_ids, pixel_values, image_sizes, position_ids, attention_mask, cache, use_recurrent_mode):
x = self.embed_new(input_ids)
cache = [None]*len(self.layers) if cache is None else cache
for i, l in enumerate(self.layers):
x, cache[i] = l(x, position_ids, attention_mask, cache[i], use_recurrent_mode)
return self.norm(x), cachedef __call__(self, input_ids, pixel_values=None, image_sizes=None, position_ids=None, attention_mask=None, cache=None, use_recurrent_mode=False):
x, cache = self.model(input_ids, pixel_values, image_sizes, position_ids, attention_mask, cache, use_recurrent_mode)
if self.untie:
return self.lm_new(x), cache
return self.model.embed_new.as_linear(x), cache
@property
def layers(self):
return self.model.layersdef __init__(self, input_dims, output_dims, r, alpha, scale, dropout, bias=False):
super().__init__()
self.linear = nn.Linear(input_dims, output_dims, bias=bias)
self.dropout = nn.Dropout(p=dropout)
self.scale = scale * (alpha / r)
scale = 1 / math.sqrt(input_dims)
self.lora_a = mx.random.uniform(low=-scale, high=scale, shape=(input_dims, r))
self.lora_b = mx.zeros(shape=(r, output_dims))
self.m = mx.linalg.norm(self._dequantized_weight(), axis=1).astype(mx.float32)
def _dequantized_weight(self):
weight = self.linear.weight
if isinstance(self.linear, nn.QuantizedLinear):
weight = mx.dequantize(weight, self.linear.scales, self.linear.biases, self.linear.group_size, self.linear.bits)
return weight
def __call__(self, x):
y = self.linear(x)
z = (self.dropout(x) @ self.lora_a) @ self.lora_b
z = y + (self.scale * z)
adapted = self._dequantized_weight() + (self.scale * self.lora_b.T) @ self.lora_a.T
denom = mx.stop_gradient(mx.linalg.norm(adapted, axis=1))
z = (self.m / denom) * z
return z.astype(x.dtype)def create_batches(data, tokenizer, batch_size, seq_length):
def _encode(x):
return [tokenizer.encode(i) for i in x]
encoded_data = [_encode(x) for x in data]
encoded_data = [x for x in encoded_data if len(x[0]+x[1]) <= seq_length+1]
if batch_size is None:
batch_size = min(len(encoded_data), 64)
else:
encoded_data = encoded_data[:(len(encoded_data) // batch_size) * batch_size]
np.random.shuffle(encoded_data)
for i in range(0, len(encoded_data), batch_size):
batch = encoded_data[i:i+batch_size]
max_len = min(max(len(q+a)-1 for q, a in batch), seq_length)
x_batch = []
y_batch = []
mask_batch = []
for q, a in batch:
combined = (q+a)[:max_len+1]
x = combined[:-1]
y = combined[1:]
pad_length = max_len - len(x)
x = x + [0] * pad_length
y = y + [0] * pad_length
mask = [False] * (len(q)-1) + [True] * (len(a)) + [False] * (pad_length)
x_batch.append(x)
y_batch.append(y)
mask_batch.append(mask)
yield mx.array(x_batch), mx.array(y_batch), mx.array(mask_batch)
def loss_fn(model, X, y, mask):
logits, _ = model(X)
logits = logits.astype(mx.float32)
ce = nn.losses.cross_entropy(logits, y, reduction='none')
masked_loss = ce * mask
return masked_loss.sum(), mask.sum()
def evaluate(model, data, tokenizer, seq_length):
model.eval()
total_loss = 0
total_samples = 0
for X, y, mask in create_batches(data, tokenizer, None, seq_length):
loss, ntoks = loss_fn(model, X, y, mask)
total_loss += loss.item()
total_samples += ntoks.item()
return total_loss / total_samples if total_samples > 0 else -1
def get_optimizer(train_data):
num_batches_per_epoch = len(list(create_batches(train_data, tokenizer, batch_size, seq_length)))
print(f'{num_batches_per_epoch=}')
num_steps = num_epochs * num_batches_per_epoch
num_warmup = num_steps // 10
max_lr, min_lr = learning_rates
if num_warmup > 2:
warmup = optim.linear_schedule(min_lr*0.1, max_lr, steps=num_warmup)
cosine = optim.cosine_decay(max_lr, num_steps - num_warmup, min_lr)
lr_schedule = optim.join_schedules([warmup, cosine], [num_warmup])
else:
lr_schedule = optim.cosine_decay(max_lr, num_steps, min_lr)
return optim.Lion(learning_rate=lr_schedule), num_steps
for arg_name in sorted(locals()):
if arg_name != 'self':
arg_value = locals()[arg_name]
if not callable(arg_value):
print(f"{arg_name}: {arg_value}")
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
print(f'--- {timestamp} ---')
train_data, eval_data = load_gsm_data(tokenizer=tokenizer)
model = load_model_for_training(lora_cfg=lora_cfg, model_cfg=model_cfg, thaws=thaws)
optimizer, num_steps = get_optimizer(train_data)
loss_and_grad_fn = nn.value_and_grad(model, loss_fn)
mx.eval(model, optimizer)
metrics = {
'steps': [],
'learning_rates': [],
'all_train_losses': [],
'avg_train_losses': [],
'val_losses': [],
'trained_toks': [],
}
step = 0
trained_toks = 0
losses = []
tic = time.perf_counter()
for epoch in range(num_epochs):
for X, y, loss_mask in create_batches(data=train_data, tokenizer=tokenizer, batch_size=batch_size, seq_length=seq_length):
model.train()
(loss, ntoks), grads = loss_and_grad_fn(model, X, y, loss_mask)
optimizer.update(model, grads)
mx.eval(loss, ntoks, model, optimizer)
losses.append(loss.item())
trained_toks += ntoks.item()
step += 1
if (step % (num_steps // 30) == 0):
avg_train_loss = np.mean(losses)
lr = optimizer.learning_rate.item()
val_loss = evaluate(model=model, data=eval_data, tokenizer=tokenizer, seq_length=seq_length)
print(f"{avg_train_loss:8.4f} ({val_loss:6.4f}) @ {step//(num_steps//30):2}/30 w/ {lr:.2e} ({time.perf_counter() - tic:.2f} sec)")
metrics['val_losses'].append(val_loss)
# print(f"{avg_train_loss:8.4f} @ {step//(num_steps//30):2}/30 w/ {lr:.2e} ({time.perf_counter() - tic:.2f} sec)")
tic = time.perf_counter()
metrics['steps'].append(step)
metrics['learning_rates'].append(lr)
metrics['all_train_losses'].extend(losses)
metrics['avg_train_losses'].append(avg_train_loss)
metrics['trained_toks'].append(trained_toks)
losses = []
trained_toks = 0
_path = f'trained_retnphi.safetensors' if model_cfg['use_retention'] else f'trained_orgnphi.safetensors'
mx.save_safetensors(_path, dict(tree_flatten(model.trainable_parameters())))
log = {
'args': {
'learning_rates': learning_rates,
'num_epochs': num_epochs,
'batch_size': batch_size,
'seq_length': seq_length,
'lora_cfg': lora_cfg,
'model_cfg': model_cfg,
'thaws': thaws,
'from_path': from_path
},
'metrics': metrics
}
with open(f'train_log_{timestamp}.json', 'w') as f:
json.dump(log, f, indent=2)
del modelmain(take=None, num_epochs=3, use_retention=False)
main(take=None, num_epochs=3, untie_embedding=False, use_retention=False)
# fire.Fire(main)