forked from IQ.Lvbs/IQ.Pilot
IQ.Pilot Prebuilt Release @ ab07000
This commit is contained in:
300
tinygrad_repo/extra/models/bert.py
Normal file
300
tinygrad_repo/extra/models/bert.py
Normal file
@@ -0,0 +1,300 @@
|
||||
import re, os
|
||||
from pathlib import Path
|
||||
from tinygrad.tensor import Tensor, cast
|
||||
from tinygrad import nn, dtypes
|
||||
from tinygrad.helpers import fetch, get_child
|
||||
from tinygrad.nn.state import get_parameters
|
||||
|
||||
# allow for monkeypatching
|
||||
Embedding = nn.Embedding
|
||||
Linear = nn.Linear
|
||||
LayerNorm = nn.LayerNorm
|
||||
|
||||
class BertForQuestionAnswering:
|
||||
def __init__(self, hidden_size=1024, intermediate_size=4096, max_position_embeddings=512, num_attention_heads=16, num_hidden_layers=24, type_vocab_size=2, vocab_size=30522, attention_probs_dropout_prob=0.1, hidden_dropout_prob=0.1):
|
||||
self.bert = Bert(hidden_size, intermediate_size, max_position_embeddings, num_attention_heads, num_hidden_layers, type_vocab_size, vocab_size, attention_probs_dropout_prob, hidden_dropout_prob)
|
||||
self.qa_outputs = Linear(hidden_size, 2)
|
||||
|
||||
def load_from_pretrained(self):
|
||||
fn = Path(__file__).parents[1] / "weights/bert_for_qa.pt"
|
||||
fetch("https://zenodo.org/record/3733896/files/model.pytorch?download=1", fn)
|
||||
fn_vocab = Path(__file__).parents[1] / "weights/bert_vocab.txt"
|
||||
fetch("https://zenodo.org/record/3733896/files/vocab.txt?download=1", fn_vocab)
|
||||
|
||||
import torch
|
||||
with open(fn, "rb") as f:
|
||||
state_dict = torch.load(f, map_location="cpu")
|
||||
|
||||
for k, v in state_dict.items():
|
||||
if "dropout" in k: continue # skip dropout
|
||||
if "pooler" in k: continue # skip pooler
|
||||
get_child(self, k).assign(v.numpy()).realize()
|
||||
|
||||
def __call__(self, input_ids:Tensor, attention_mask:Tensor, token_type_ids:Tensor):
|
||||
sequence_output = self.bert(input_ids, attention_mask, token_type_ids)
|
||||
logits = self.qa_outputs(sequence_output)
|
||||
start_logits, end_logits = logits.chunk(2, dim=-1)
|
||||
start_logits = start_logits.reshape(-1, 1)
|
||||
end_logits = end_logits.reshape(-1, 1)
|
||||
|
||||
return Tensor.stack(start_logits, end_logits)
|
||||
|
||||
class BertForPretraining:
|
||||
def __init__(self, hidden_size:int=1024, intermediate_size:int=4096, max_position_embeddings:int=512, num_attention_heads:int=16, num_hidden_layers:int=24, type_vocab_size:int=2, vocab_size:int=30522, attention_probs_dropout_prob:float=0.1, hidden_dropout_prob:float=0.1):
|
||||
"""Default is BERT-large"""
|
||||
self.bert = Bert(hidden_size, intermediate_size, max_position_embeddings, num_attention_heads, num_hidden_layers, type_vocab_size, vocab_size, attention_probs_dropout_prob, hidden_dropout_prob)
|
||||
self.cls = BertPreTrainingHeads(hidden_size, vocab_size, self.bert.embeddings.word_embeddings.weight)
|
||||
|
||||
def __call__(self, input_ids:Tensor, attention_mask:Tensor, masked_lm_positions:Tensor, token_type_ids:Tensor):
|
||||
output = self.bert(input_ids, attention_mask, token_type_ids)
|
||||
return self.cls(output, masked_lm_positions)
|
||||
|
||||
# Reference has residual on denominator: https://github.com/mlcommons/training/blob/master/language_model/tensorflow/bert/run_pretraining.py#L315
|
||||
def sparse_categorical_crossentropy(self, predictions:Tensor, labels:Tensor, ignore_index=-1):
|
||||
log_probs, loss_mask = predictions.log_softmax(dtype=dtypes.float), (labels != ignore_index)
|
||||
y_counter = Tensor.arange(predictions.shape[-1], device=predictions.device).unsqueeze(0).expand(labels.numel(), predictions.shape[-1])
|
||||
y = ((y_counter == labels.flatten().reshape(-1, 1)) * loss_mask.reshape(-1, 1)).reshape(*labels.shape, predictions.shape[-1])
|
||||
return -((log_probs * y).sum()) / (loss_mask.sum() + 1e-5) # Small constant to avoid division by zero
|
||||
|
||||
def loss(self, prediction_logits:Tensor, seq_relationship_logits:Tensor, masked_lm_ids:Tensor, masked_lm_weights:Tensor, next_sentence_labels:Tensor):
|
||||
masked_lm_loss = self.sparse_categorical_crossentropy(prediction_logits, masked_lm_ids, ignore_index=masked_lm_weights)
|
||||
next_sentence_loss = seq_relationship_logits.binary_crossentropy_logits(next_sentence_labels)
|
||||
return masked_lm_loss + next_sentence_loss
|
||||
|
||||
def accuracy(self, prediction_logits:Tensor, seq_relationship_logits:Tensor, masked_lm_ids:Tensor, masked_lm_weights:Tensor, next_sentence_labels:Tensor):
|
||||
valid = masked_lm_ids != 0
|
||||
masked_lm_predictions = prediction_logits.argmax(-1)
|
||||
masked_lm_correct = (masked_lm_predictions == masked_lm_ids) * valid
|
||||
masked_lm_loss = self.sparse_categorical_crossentropy(prediction_logits, masked_lm_ids, ignore_index=masked_lm_weights)
|
||||
|
||||
seq_relationship_predictions = seq_relationship_logits.argmax(-1)
|
||||
seq_relationship_correct = (seq_relationship_predictions == next_sentence_labels)
|
||||
next_sentence_loss = seq_relationship_logits.binary_crossentropy_logits(next_sentence_labels)
|
||||
|
||||
# NOTE: .float().sum() to prevent overflow with large BS since default acc of bool is in default_float
|
||||
# TODO: is it okay that next_sentence_loss is half here?
|
||||
return masked_lm_correct.float().sum() / valid.float().sum(), seq_relationship_correct.float().mean(), masked_lm_loss, next_sentence_loss.float()
|
||||
|
||||
def load_from_pretrained(self, tf_weight_path:str=Path(__file__).parent.parent / "datasets" / "wiki"):
|
||||
os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2' # Mute tf flag info
|
||||
# load from tensorflow
|
||||
import tensorflow as tf
|
||||
import numpy as np
|
||||
|
||||
state_dict = {}
|
||||
for name, _ in tf.train.list_variables(str(tf_weight_path)):
|
||||
state_dict[name] = tf.train.load_variable(str(tf_weight_path), name)
|
||||
|
||||
for k, v in state_dict.items():
|
||||
m = k.split("/")
|
||||
if any(n in ["adam_v", "adam_m", "global_step", "LAMB", "LAMB_1", "beta1_power", "beta2_power"] for n in m):
|
||||
continue
|
||||
|
||||
pointer = self
|
||||
n = m[-1] # this is just to stop python from complaining about possibly unbound local variable
|
||||
for i, n in enumerate(m):
|
||||
if re.fullmatch(r'[A-Za-z]+_\d+', n):
|
||||
l = re.split(r'_(\d+)', n)[:-1]
|
||||
else:
|
||||
l = [n]
|
||||
if l[0] in ["kernel", "gamma", "output_weights"]:
|
||||
pointer = getattr(pointer, "weight")
|
||||
elif l[0] in ["output_bias", "beta"]:
|
||||
pointer = getattr(pointer, "bias")
|
||||
elif l[0] == "pooler":
|
||||
pointer = getattr(getattr(self, "cls"), "pooler")
|
||||
else:
|
||||
pointer = getattr(pointer, l[0])
|
||||
if len(l) == 2: # layers
|
||||
pointer = pointer[int(l[1])]
|
||||
if n[-11:] == "_embeddings":
|
||||
pointer = getattr(pointer, "weight")
|
||||
elif n == "kernel":
|
||||
v = np.transpose(v)
|
||||
cast(Tensor, pointer).assign(v).realize()
|
||||
|
||||
params = get_parameters(self)
|
||||
count = 0
|
||||
for p in params:
|
||||
param_count = 1
|
||||
for s in p.shape:
|
||||
param_count *= s
|
||||
count += param_count
|
||||
print(f"Total parameters: {count / 1000 / 1000}M")
|
||||
return self
|
||||
|
||||
class BertPreTrainingHeads:
|
||||
def __init__(self, hidden_size:int, vocab_size:int, embeddings_weight:Tensor):
|
||||
self.predictions = BertLMPredictionHead(hidden_size, vocab_size, embeddings_weight)
|
||||
self.pooler = BertPooler(hidden_size)
|
||||
self.seq_relationship = Linear(hidden_size, 2)
|
||||
|
||||
def __call__(self, sequence_output:Tensor, masked_lm_positions:Tensor):
|
||||
prediction_logits = self.predictions(gather(sequence_output, masked_lm_positions))
|
||||
seq_relationship_logits = self.seq_relationship(self.pooler(sequence_output))
|
||||
return prediction_logits, seq_relationship_logits
|
||||
|
||||
class BertLMPredictionHead:
|
||||
def __init__(self, hidden_size:int, vocab_size:int, embeddings_weight:Tensor):
|
||||
self.transform = BertPredictionHeadTransform(hidden_size)
|
||||
self.embedding_weight = embeddings_weight
|
||||
self.bias = Tensor.zeros(vocab_size, dtype=dtypes.float32)
|
||||
|
||||
def __call__(self, hidden_states:Tensor):
|
||||
return self.transform(hidden_states) @ self.embedding_weight.T + self.bias
|
||||
|
||||
class BertPredictionHeadTransform:
|
||||
def __init__(self, hidden_size:int):
|
||||
self.dense = Linear(hidden_size, hidden_size)
|
||||
self.LayerNorm = LayerNorm(hidden_size, eps=1e-12)
|
||||
|
||||
def __call__(self, hidden_states:Tensor):
|
||||
return self.LayerNorm(gelu(self.dense(hidden_states)))
|
||||
|
||||
class BertPooler:
|
||||
def __init__(self, hidden_size:int):
|
||||
self.dense = Linear(hidden_size, hidden_size)
|
||||
|
||||
def __call__(self, hidden_states:Tensor):
|
||||
return self.dense(hidden_states[:, 0]).tanh()
|
||||
|
||||
def gather(prediction_logits:Tensor, masked_lm_positions:Tensor):
|
||||
counter = Tensor.arange(prediction_logits.shape[1], device=prediction_logits.device).reshape(1, 1, prediction_logits.shape[1]).expand(*masked_lm_positions.shape, prediction_logits.shape[1])
|
||||
onehot = counter == masked_lm_positions.unsqueeze(2).expand(*masked_lm_positions.shape, prediction_logits.shape[1])
|
||||
return onehot @ prediction_logits
|
||||
|
||||
class Bert:
|
||||
def __init__(self, hidden_size, intermediate_size, max_position_embeddings, num_attention_heads, num_hidden_layers, type_vocab_size, vocab_size, attention_probs_dropout_prob, hidden_dropout_prob):
|
||||
self.embeddings = BertEmbeddings(hidden_size, max_position_embeddings, type_vocab_size, vocab_size, hidden_dropout_prob)
|
||||
self.encoder = BertEncoder(hidden_size, intermediate_size, num_attention_heads, num_hidden_layers, attention_probs_dropout_prob, hidden_dropout_prob)
|
||||
|
||||
def __call__(self, input_ids, attention_mask, token_type_ids):
|
||||
extended_attention_mask = attention_mask.unsqueeze(1).unsqueeze(2)
|
||||
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
|
||||
|
||||
embedding_output = self.embeddings(input_ids, token_type_ids)
|
||||
encoder_outputs = self.encoder(embedding_output, extended_attention_mask)
|
||||
|
||||
return encoder_outputs
|
||||
|
||||
class BertEmbeddings:
|
||||
def __init__(self, hidden_size, max_position_embeddings, type_vocab_size, vocab_size, hidden_dropout_prob):
|
||||
self.word_embeddings = Embedding(vocab_size, hidden_size)
|
||||
self.position_embeddings = Embedding(max_position_embeddings, hidden_size)
|
||||
self.token_type_embeddings = Embedding(type_vocab_size, hidden_size)
|
||||
self.LayerNorm = LayerNorm(hidden_size, eps=1e-12)
|
||||
self.dropout = hidden_dropout_prob
|
||||
|
||||
def __call__(self, input_ids, token_type_ids):
|
||||
input_shape = input_ids.shape
|
||||
seq_length = input_shape[1]
|
||||
|
||||
position_ids = Tensor.arange(seq_length, device=input_ids.device).unsqueeze(0).expand(*input_shape)
|
||||
words_embeddings = self.word_embeddings(input_ids)
|
||||
position_embeddings = self.position_embeddings(position_ids)
|
||||
token_type_embeddings = self.token_type_embeddings(token_type_ids)
|
||||
|
||||
embeddings = words_embeddings + position_embeddings + token_type_embeddings
|
||||
embeddings = self.LayerNorm(embeddings)
|
||||
embeddings = embeddings.dropout(self.dropout)
|
||||
return embeddings
|
||||
|
||||
class BertEncoder:
|
||||
def __init__(self, hidden_size, intermediate_size, num_attention_heads, num_hidden_layers, attention_probs_dropout_prob, hidden_dropout_prob):
|
||||
self.layer = [BertLayer(hidden_size, intermediate_size, num_attention_heads, attention_probs_dropout_prob, hidden_dropout_prob) for _ in range(num_hidden_layers)]
|
||||
|
||||
def __call__(self, hidden_states, attention_mask):
|
||||
for layer in self.layer:
|
||||
hidden_states = layer(hidden_states, attention_mask)
|
||||
return hidden_states
|
||||
|
||||
class BertLayer:
|
||||
def __init__(self, hidden_size, intermediate_size, num_attention_heads, attention_probs_dropout_prob, hidden_dropout_prob):
|
||||
self.attention = BertAttention(hidden_size, num_attention_heads, attention_probs_dropout_prob, hidden_dropout_prob)
|
||||
self.intermediate = BertIntermediate(hidden_size, intermediate_size)
|
||||
self.output = BertOutput(hidden_size, intermediate_size, hidden_dropout_prob)
|
||||
|
||||
def __call__(self, hidden_states, attention_mask):
|
||||
attention_output = self.attention(hidden_states, attention_mask)
|
||||
intermediate_output = self.intermediate(attention_output)
|
||||
layer_output = self.output(intermediate_output, attention_output)
|
||||
return layer_output
|
||||
|
||||
class BertOutput:
|
||||
def __init__(self, hidden_size, intermediate_size, hidden_dropout_prob):
|
||||
self.dense = Linear(intermediate_size, hidden_size)
|
||||
self.LayerNorm = LayerNorm(hidden_size, eps=1e-12)
|
||||
self.dropout = hidden_dropout_prob
|
||||
|
||||
def __call__(self, hidden_states, input_tensor):
|
||||
hidden_states = self.dense(hidden_states)
|
||||
hidden_states = hidden_states.dropout(self.dropout)
|
||||
hidden_states = self.LayerNorm(hidden_states + input_tensor)
|
||||
return hidden_states
|
||||
|
||||
def gelu(x):
|
||||
return x * 0.5 * (1.0 + (x / 1.41421).erf())
|
||||
|
||||
class BertIntermediate:
|
||||
def __init__(self, hidden_size, intermediate_size):
|
||||
self.dense = Linear(hidden_size, intermediate_size)
|
||||
|
||||
def __call__(self, hidden_states):
|
||||
x = self.dense(hidden_states)
|
||||
# tinygrad gelu is openai gelu but we need the original bert gelu
|
||||
# NOTE: contiguous for speed
|
||||
return gelu(x).contiguous()
|
||||
|
||||
class BertAttention:
|
||||
def __init__(self, hidden_size, num_attention_heads, attention_probs_dropout_prob, hidden_dropout_prob):
|
||||
self.self = BertSelfAttention(hidden_size, num_attention_heads, attention_probs_dropout_prob)
|
||||
self.output = BertSelfOutput(hidden_size, hidden_dropout_prob)
|
||||
|
||||
def __call__(self, hidden_states, attention_mask):
|
||||
self_output = self.self(hidden_states, attention_mask)
|
||||
attention_output = self.output(self_output, hidden_states)
|
||||
return attention_output
|
||||
|
||||
class BertSelfAttention:
|
||||
def __init__(self, hidden_size, num_attention_heads, attention_probs_dropout_prob):
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.attention_head_size = int(hidden_size / num_attention_heads)
|
||||
self.all_head_size = self.num_attention_heads * self.attention_head_size
|
||||
|
||||
self.query = Linear(hidden_size, self.all_head_size)
|
||||
self.key = Linear(hidden_size, self.all_head_size)
|
||||
self.value = Linear(hidden_size, self.all_head_size)
|
||||
|
||||
self.dropout = attention_probs_dropout_prob
|
||||
|
||||
def __call__(self, hidden_states, attention_mask):
|
||||
mixed_query_layer = self.query(hidden_states)
|
||||
mixed_key_layer = self.key(hidden_states)
|
||||
mixed_value_layer = self.value(hidden_states)
|
||||
|
||||
query_layer = self.transpose_for_scores(mixed_query_layer)
|
||||
key_layer = self.transpose_for_scores(mixed_key_layer)
|
||||
value_layer = self.transpose_for_scores(mixed_value_layer)
|
||||
|
||||
context_layer = Tensor.scaled_dot_product_attention(query_layer, key_layer, value_layer, attention_mask, self.dropout)
|
||||
|
||||
context_layer = context_layer.transpose(1, 2)
|
||||
context_layer = context_layer.reshape(context_layer.shape[0], context_layer.shape[1], self.all_head_size)
|
||||
|
||||
return context_layer
|
||||
|
||||
def transpose_for_scores(self, x):
|
||||
x = x.reshape(x.shape[0], x.shape[1], self.num_attention_heads, self.attention_head_size)
|
||||
return x.transpose(1, 2)
|
||||
|
||||
class BertSelfOutput:
|
||||
def __init__(self, hidden_size, hidden_dropout_prob):
|
||||
self.dense = Linear(hidden_size, hidden_size)
|
||||
self.LayerNorm = LayerNorm(hidden_size, eps=1e-12)
|
||||
self.dropout = hidden_dropout_prob
|
||||
|
||||
def __call__(self, hidden_states, input_tensor):
|
||||
hidden_states = self.dense(hidden_states)
|
||||
hidden_states = hidden_states.dropout(self.dropout)
|
||||
hidden_states = self.LayerNorm(hidden_states + input_tensor)
|
||||
return hidden_states
|
||||
480
tinygrad_repo/extra/models/clip.py
Normal file
480
tinygrad_repo/extra/models/clip.py
Normal file
@@ -0,0 +1,480 @@
|
||||
from tinygrad import Tensor, dtypes
|
||||
from tinygrad.helpers import fetch
|
||||
from tinygrad.nn import Linear, LayerNorm, Embedding, Conv2d
|
||||
|
||||
from typing import List, Optional, Union, Tuple, Dict
|
||||
from abc import ABC, abstractmethod
|
||||
from functools import lru_cache
|
||||
import numpy as np
|
||||
import re, gzip
|
||||
|
||||
# Allow for monkeypatching for mlperf.
|
||||
gelu = Tensor.gelu
|
||||
|
||||
@lru_cache()
|
||||
def default_bpe():
|
||||
# Clip tokenizer, taken from https://github.com/openai/CLIP/blob/main/clip/simple_tokenizer.py (MIT license)
|
||||
return fetch("https://github.com/openai/CLIP/raw/main/clip/bpe_simple_vocab_16e6.txt.gz", "bpe_simple_vocab_16e6.txt.gz")
|
||||
|
||||
class Tokenizer:
|
||||
"""
|
||||
Namespace for CLIP Text Tokenizer components.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def get_pairs(word):
|
||||
"""
|
||||
Return set of symbol pairs in a word.
|
||||
Word is represented as tuple of symbols (symbols being variable-length strings).
|
||||
"""
|
||||
return set(zip(word, word[1:]))
|
||||
@staticmethod
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
@staticmethod
|
||||
def bytes_to_unicode():
|
||||
"""
|
||||
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
||||
The reversible bpe codes work on unicode strings.
|
||||
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
||||
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
||||
This is a significant percentage of your normal, say, 32K bpe vocab.
|
||||
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
||||
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
||||
"""
|
||||
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8+n)
|
||||
n += 1
|
||||
cs = [chr(n) for n in cs]
|
||||
return dict(zip(bs, cs))
|
||||
class ClipTokenizer:
|
||||
def __init__(self, version=None):
|
||||
self.byte_encoder, self.version = Tokenizer.bytes_to_unicode(), version
|
||||
merges = gzip.open(default_bpe()).read().decode("utf-8").split('\n')
|
||||
merges = merges[1:49152-256-2+1]
|
||||
merges = [tuple(merge.split()) for merge in merges]
|
||||
vocab = list(Tokenizer.bytes_to_unicode().values())
|
||||
vocab = vocab + [v+'</w>' for v in vocab]
|
||||
for merge in merges:
|
||||
vocab.append(''.join(merge))
|
||||
if self.version == "sd_mlperf_v5_0":
|
||||
import regex
|
||||
vocab.extend(['<start_of_text>', '<end_of_text>'])
|
||||
self.cache = {'<start_of_text>': '<start_of_text>', '<end_of_text>': '<end_of_text>'}
|
||||
self.pat = regex.compile(r"""<start_of_text>|<end_of_text>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", regex.IGNORECASE)
|
||||
else:
|
||||
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
|
||||
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
|
||||
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[^\s]+""", re.IGNORECASE)
|
||||
self.encoder = dict(zip(vocab, range(len(vocab))))
|
||||
self.bpe_ranks = dict(zip(merges, range(len(merges))))
|
||||
|
||||
def bpe(self, token):
|
||||
if token in self.cache:
|
||||
return self.cache[token]
|
||||
word = tuple(token[:-1]) + ( token[-1] + '</w>',)
|
||||
pairs = Tokenizer.get_pairs(word)
|
||||
|
||||
if not pairs:
|
||||
return token+'</w>'
|
||||
|
||||
while True:
|
||||
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
||||
if bigram not in self.bpe_ranks:
|
||||
break
|
||||
first, second = bigram
|
||||
new_word = []
|
||||
i = 0
|
||||
while i < len(word):
|
||||
try:
|
||||
j = word.index(first, i)
|
||||
new_word.extend(word[i:j])
|
||||
i = j
|
||||
except Exception:
|
||||
new_word.extend(word[i:])
|
||||
break
|
||||
|
||||
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
||||
new_word.append(first+second)
|
||||
i += 2
|
||||
else:
|
||||
new_word.append(word[i])
|
||||
i += 1
|
||||
new_word = tuple(new_word)
|
||||
word = new_word
|
||||
if len(word) == 1:
|
||||
break
|
||||
pairs = Tokenizer.get_pairs(word)
|
||||
word = ' '.join(word)
|
||||
self.cache[token] = word
|
||||
return word
|
||||
|
||||
def encode(self, text:str, pad_with_zeros:bool=False) -> List[int]:
|
||||
bpe_tokens: List[int] = []
|
||||
if self.version == "sd_mlperf_v5_0":
|
||||
import regex, ftfy, html
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text)).strip()
|
||||
text = Tokenizer.whitespace_clean(text).lower()
|
||||
re_module = regex
|
||||
else:
|
||||
text = Tokenizer.whitespace_clean(text.strip()).lower()
|
||||
re_module = re
|
||||
|
||||
for token in re_module.findall(self.pat, text):
|
||||
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
||||
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
||||
# Truncation, keeping two slots for start and end tokens.
|
||||
if len(bpe_tokens) > 75:
|
||||
bpe_tokens = bpe_tokens[:75]
|
||||
return [49406] + bpe_tokens + [49407] + ([0] if pad_with_zeros else [49407]) * (77 - len(bpe_tokens) - 2)
|
||||
|
||||
|
||||
class Embedder(ABC):
|
||||
input_key: str
|
||||
@abstractmethod
|
||||
def __call__(self, x:Union[str,List[str],Tensor]) -> Union[Tensor,Tuple[Tensor,...]]:
|
||||
pass
|
||||
|
||||
|
||||
class Closed:
|
||||
"""
|
||||
Namespace for OpenAI CLIP model components.
|
||||
"""
|
||||
class ClipMlp:
|
||||
def __init__(self):
|
||||
self.fc1 = Linear(768, 3072)
|
||||
self.fc2 = Linear(3072, 768)
|
||||
|
||||
def __call__(self, h:Tensor) -> Tensor:
|
||||
h = self.fc1(h)
|
||||
h = h.quick_gelu()
|
||||
h = self.fc2(h)
|
||||
return h
|
||||
|
||||
class ClipAttention:
|
||||
def __init__(self):
|
||||
self.embed_dim = 768
|
||||
self.num_heads = 12
|
||||
self.head_dim = self.embed_dim // self.num_heads
|
||||
self.k_proj = Linear(self.embed_dim, self.embed_dim)
|
||||
self.v_proj = Linear(self.embed_dim, self.embed_dim)
|
||||
self.q_proj = Linear(self.embed_dim, self.embed_dim)
|
||||
self.out_proj = Linear(self.embed_dim, self.embed_dim)
|
||||
|
||||
def __call__(self, hidden_states:Tensor, causal_attention_mask:Tensor) -> Tensor:
|
||||
bsz, tgt_len, embed_dim = hidden_states.shape
|
||||
q,k,v = self.q_proj(hidden_states), self.k_proj(hidden_states), self.v_proj(hidden_states)
|
||||
q,k,v = [x.reshape(bsz, tgt_len, self.num_heads, self.head_dim).transpose(1, 2) for x in (q,k,v)]
|
||||
attn_output = Tensor.scaled_dot_product_attention(q, k, v, attn_mask=causal_attention_mask)
|
||||
return self.out_proj(attn_output.transpose(1, 2).reshape(bsz, tgt_len, embed_dim))
|
||||
|
||||
class ClipEncoderLayer:
|
||||
def __init__(self):
|
||||
self.self_attn = Closed.ClipAttention()
|
||||
self.layer_norm1 = LayerNorm(768)
|
||||
self.mlp = Closed.ClipMlp()
|
||||
self.layer_norm2 = LayerNorm(768)
|
||||
|
||||
def __call__(self, hidden_states:Tensor, causal_attention_mask:Tensor) -> Tensor:
|
||||
residual = hidden_states
|
||||
hidden_states = self.layer_norm1(hidden_states)
|
||||
hidden_states = self.self_attn(hidden_states, causal_attention_mask)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.layer_norm2(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
return hidden_states
|
||||
|
||||
class ClipTextEmbeddings:
|
||||
def __init__(self):
|
||||
self.token_embedding = Embedding(49408, 768)
|
||||
self.position_embedding = Embedding(77, 768)
|
||||
|
||||
def __call__(self, input_ids:Tensor, position_ids:Tensor) -> Tensor:
|
||||
return self.token_embedding(input_ids) + self.position_embedding(position_ids)
|
||||
|
||||
class ClipEncoder:
|
||||
def __init__(self, layer_count:int=12):
|
||||
self.layers = [Closed.ClipEncoderLayer() for _ in range(layer_count)]
|
||||
|
||||
def __call__(self, x:Tensor, causal_attention_mask:Tensor, ret_layer_idx:Optional[int]=None) -> Tensor:
|
||||
# the indexing of layers is NOT off by 1, the original code considers the "input" as the first hidden state
|
||||
layers = self.layers if ret_layer_idx is None else self.layers[:ret_layer_idx]
|
||||
for l in layers:
|
||||
x = l(x, causal_attention_mask)
|
||||
return x
|
||||
|
||||
class ClipTextTransformer:
|
||||
def __init__(self, ret_layer_idx:Optional[int]=None):
|
||||
self.embeddings = Closed.ClipTextEmbeddings()
|
||||
self.encoder = Closed.ClipEncoder()
|
||||
self.final_layer_norm = LayerNorm(768)
|
||||
self.ret_layer_idx = ret_layer_idx
|
||||
|
||||
def __call__(self, input_ids:Tensor) -> Tensor:
|
||||
x = self.embeddings(input_ids, Tensor.arange(input_ids.shape[1]).reshape(1, -1))
|
||||
x = self.encoder(x, Tensor.full((1, 1, 77, 77), float("-inf")).triu(1), self.ret_layer_idx)
|
||||
return self.final_layer_norm(x) if (self.ret_layer_idx is None) else x
|
||||
|
||||
class ClipTextModel:
|
||||
def __init__(self, ret_layer_idx:Optional[int]):
|
||||
self.text_model = Closed.ClipTextTransformer(ret_layer_idx=ret_layer_idx)
|
||||
|
||||
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/encoders/modules.py#L331
|
||||
class FrozenClosedClipEmbedder(Embedder):
|
||||
def __init__(self, ret_layer_idx:Optional[int]=None):
|
||||
self.tokenizer = Tokenizer.ClipTokenizer()
|
||||
self.transformer = Closed.ClipTextModel(ret_layer_idx)
|
||||
self.input_key = "txt"
|
||||
|
||||
def __call__(self, texts:Union[str,List[str],Tensor]) -> Union[Tensor,Tuple[Tensor,...]]:
|
||||
if isinstance(texts, str): texts = [texts]
|
||||
assert isinstance(texts, (list,tuple)), f"expected list of strings, got {type(texts).__name__}"
|
||||
tokens = Tensor.cat(*[Tensor(self.tokenizer.encode(text)) for text in texts], dim=0)
|
||||
return self.transformer.text_model(tokens.reshape(len(texts),-1))
|
||||
|
||||
|
||||
class Open:
|
||||
"""
|
||||
Namespace for OpenCLIP model components.
|
||||
"""
|
||||
class MultiheadAttention:
|
||||
def __init__(self, dims:int, n_heads:int):
|
||||
self.dims = dims
|
||||
self.n_heads = n_heads
|
||||
self.d_head = self.dims // self.n_heads
|
||||
|
||||
self.in_proj_bias = Tensor.empty(3*dims)
|
||||
self.in_proj_weight = Tensor.empty(3*dims, dims)
|
||||
self.out_proj = Linear(dims, dims)
|
||||
|
||||
def __call__(self, x:Tensor, attn_mask:Optional[Tensor]=None) -> Tensor:
|
||||
T,B,C = x.shape
|
||||
|
||||
proj = x.linear(self.in_proj_weight.T, self.in_proj_bias)
|
||||
proj = proj.unflatten(-1, (3,C)).unsqueeze(0).transpose(0, -2)
|
||||
|
||||
q,k,v = [y.reshape(T, B*self.n_heads, self.d_head).transpose(0, 1).reshape(B, self.n_heads, T, self.d_head) for y in proj.chunk(3)]
|
||||
|
||||
attn_output = Tensor.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
||||
attn_output = attn_output.permute(2, 0, 1, 3).reshape(T, B, C)
|
||||
attn_output = self.out_proj(attn_output)
|
||||
|
||||
return attn_output
|
||||
|
||||
class Mlp:
|
||||
def __init__(self, dims, hidden_dims):
|
||||
self.c_fc = Linear(dims, hidden_dims)
|
||||
self.c_proj = Linear(hidden_dims, dims)
|
||||
self.gelu = gelu
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return x.sequential([self.c_fc, self.gelu, self.c_proj])
|
||||
|
||||
# https://github.com/mlfoundations/open_clip/blob/58e4e39aaabc6040839b0d2a7e8bf20979e4558a/src/open_clip/transformer.py#L210
|
||||
class ResidualAttentionBlock:
|
||||
def __init__(self, dims:int, n_heads:int, mlp_ratio:float):
|
||||
self.ln_1 = LayerNorm(dims)
|
||||
self.attn = Open.MultiheadAttention(dims, n_heads)
|
||||
|
||||
self.ln_2 = LayerNorm(dims)
|
||||
self.mlp = Open.Mlp(dims, int(dims * mlp_ratio))
|
||||
|
||||
def __call__(self, x:Tensor, attn_mask:Optional[Tensor]=None, transpose:bool=False) -> Tensor:
|
||||
q_x = self.ln_1(x)
|
||||
attn_out = self.attn(q_x.transpose(0, 1) if transpose else q_x, attn_mask=attn_mask)
|
||||
attn_out = attn_out.transpose(0, 1) if transpose else attn_out
|
||||
x = x + attn_out
|
||||
x = x + self.mlp(self.ln_2(x))
|
||||
return x
|
||||
|
||||
# https://github.com/mlfoundations/open_clip/blob/58e4e39aaabc6040839b0d2a7e8bf20979e4558a/src/open_clip/transformer.py#L317
|
||||
class ClipTransformer:
|
||||
def __init__(self, dims:int, layers:int, n_heads:int, mlp_ratio:float=4.0):
|
||||
self.resblocks = [
|
||||
Open.ResidualAttentionBlock(dims, n_heads, mlp_ratio) for _ in range(layers)
|
||||
]
|
||||
|
||||
def __call__(self, x:Tensor, attn_mask:Optional[Tensor]=None) -> Tensor:
|
||||
for r in self.resblocks:
|
||||
x = r(x, attn_mask=attn_mask, transpose=True)
|
||||
return x
|
||||
|
||||
# https://github.com/mlfoundations/open_clip/blob/58e4e39aaabc6040839b0d2a7e8bf20979e4558a/src/open_clip/model.py#L220
|
||||
# https://github.com/mlfoundations/open_clip/blob/58e4e39aaabc6040839b0d2a7e8bf20979e4558a/src/open_clip/transformer.py#L661
|
||||
class ClipTextTransformer:
|
||||
def __init__(self, width:int, n_heads:int, layers:int, vocab_size:int=49408, ctx_length:int=77):
|
||||
self.token_embedding = Embedding(vocab_size, width)
|
||||
self.positional_embedding = Tensor.empty(ctx_length, width)
|
||||
self.transformer = Open.ClipTransformer(width, layers, n_heads)
|
||||
self.ln_final = LayerNorm(width)
|
||||
self.text_projection = Tensor.empty(width, width)
|
||||
self.attn_mask = Tensor.full((77, 77), float("-inf")).triu(1).realize()
|
||||
|
||||
def __call__(self, text:Tensor) -> Tensor:
|
||||
seq_len = text.shape[1]
|
||||
|
||||
x = self.token_embedding(text)
|
||||
x = x + self.positional_embedding[:seq_len]
|
||||
x = self.transformer(x, attn_mask=self.attn_mask)
|
||||
x = self.ln_final(x)
|
||||
|
||||
pooled = x[:, text.argmax(dim=-1)] @ self.text_projection
|
||||
return pooled
|
||||
|
||||
class ClipVisionTransformer:
|
||||
def __init__(self, width:int, layers:int, d_head:int, image_size:int, patch_size:int):
|
||||
grid_size = image_size // patch_size
|
||||
n_heads = width // d_head
|
||||
assert n_heads * d_head == width
|
||||
|
||||
self.conv1 = Conv2d(3, width, kernel_size=patch_size, stride=patch_size, bias=False)
|
||||
|
||||
self.class_embedding = Tensor.empty(width)
|
||||
self.positional_embedding = Tensor.empty(grid_size * grid_size + 1, width)
|
||||
self.transformer = Open.ClipTransformer(width, layers, n_heads)
|
||||
self.ln_pre = LayerNorm(width)
|
||||
self.ln_post = LayerNorm(width)
|
||||
self.proj = Tensor.empty(width, 1024)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
x = self.conv1(x)
|
||||
x = x.reshape(x.shape[0], x.shape[1], -1).permute(0, 2, 1)
|
||||
x = self.class_embedding.reshape(1, 1, -1).expand(x.shape[0], 1, -1).cat(x, dim=1)
|
||||
x = x + self.positional_embedding
|
||||
|
||||
x = self.ln_pre(x)
|
||||
x = self.transformer(x)
|
||||
x = self.ln_post(x)
|
||||
|
||||
pooled = x[:, 0] @ self.proj
|
||||
return pooled
|
||||
|
||||
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/encoders/modules.py#L396
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/encoders/modules.py#L498
|
||||
class FrozenOpenClipEmbedder(Embedder):
|
||||
def __init__(self, dims:int, n_heads:int, layers:int, return_pooled:bool, ln_penultimate:bool=False, clip_tokenizer_version=None):
|
||||
self.tokenizer = Tokenizer.ClipTokenizer(version=clip_tokenizer_version)
|
||||
self.model = Open.ClipTextTransformer(dims, n_heads, layers)
|
||||
self.return_pooled = return_pooled
|
||||
self.input_key = "txt"
|
||||
self.ln_penultimate = ln_penultimate
|
||||
|
||||
def tokenize(self, text:str, device:Optional[str]=None) -> Tensor:
|
||||
return Tensor(self.tokenizer.encode(text, pad_with_zeros=True), dtype=dtypes.int32, device=device).reshape(1,-1)
|
||||
|
||||
def text_transformer_forward(self, x:Tensor, attn_mask:Optional[Tensor]=None):
|
||||
for r in self.model.transformer.resblocks:
|
||||
x, penultimate = r(x, attn_mask=attn_mask), x
|
||||
return x.permute(1, 0, 2), penultimate.permute(1, 0, 2)
|
||||
|
||||
def embed_tokens(self, tokens:Tensor) -> Union[Tensor,Tuple[Tensor,...]]:
|
||||
x = self.model.token_embedding(tokens).add(self.model.positional_embedding).permute(1,0,2)
|
||||
x, penultimate = self.text_transformer_forward(x, attn_mask=self.model.attn_mask)
|
||||
|
||||
if self.ln_penultimate:
|
||||
penultimate = self.model.ln_final(penultimate)
|
||||
|
||||
if self.return_pooled:
|
||||
x = self.model.ln_final(x)
|
||||
index = tokens.argmax(axis=-1).reshape(-1,1,1).expand(x.shape[0],1,x.shape[-1])
|
||||
pooled = x.gather(1, index).squeeze(1) @ self.model.text_projection
|
||||
return penultimate, pooled
|
||||
else:
|
||||
return penultimate
|
||||
|
||||
def __call__(self, texts:Union[str,List[str],Tensor]) -> Union[Tensor,Tuple[Tensor,...]]:
|
||||
if isinstance(texts, str): texts = [texts]
|
||||
assert isinstance(texts, (list,tuple)), f"expected list of strings, got {type(texts).__name__}"
|
||||
tokens = Tensor.cat(*[self.tokenize(text) for text in texts], dim=0)
|
||||
return self.embed_tokens(tokens)
|
||||
|
||||
|
||||
clip_configs: Dict = {
|
||||
"ViT-H-14": {
|
||||
"dims": 1024,
|
||||
"vision_cfg": {
|
||||
"width": 1280,
|
||||
"layers": 32,
|
||||
"d_head": 80,
|
||||
"image_size": 224,
|
||||
"patch_size": 14,
|
||||
},
|
||||
"text_cfg": {
|
||||
"width": 1024,
|
||||
"n_heads": 16,
|
||||
"layers": 24,
|
||||
"ctx_length": 77,
|
||||
"vocab_size": 49408,
|
||||
},
|
||||
"return_pooled": False,
|
||||
"ln_penultimate": True,
|
||||
}
|
||||
}
|
||||
|
||||
class OpenClipEncoder:
|
||||
def __init__(self, dims:int, text_cfg:Dict, vision_cfg:Dict, **_):
|
||||
self.visual = Open.ClipVisionTransformer(**vision_cfg)
|
||||
|
||||
text = Open.ClipTextTransformer(**text_cfg)
|
||||
self.transformer = text.transformer
|
||||
self.token_embedding = text.token_embedding
|
||||
self.positional_embedding = text.positional_embedding
|
||||
self.ln_final = text.ln_final
|
||||
self.text_projection = text.text_projection
|
||||
|
||||
self.attn_mask = Tensor.full((77, 77), float("-inf")).triu(1).realize()
|
||||
self.mean = Tensor([0.48145466, 0.45782750, 0.40821073]).reshape(-1, 1, 1)
|
||||
self.std = Tensor([0.26862954, 0.26130258, 0.27577711]).reshape(-1, 1, 1)
|
||||
|
||||
# TODO:
|
||||
# Should be doable in pure tinygrad, would just require some work and verification.
|
||||
# This is very desirable since it would allow for full generation->evaluation in a single JIT call.
|
||||
def prepare_image(self, image) -> Tensor:
|
||||
from PIL import Image
|
||||
SIZE = 224
|
||||
w, h = image.size
|
||||
scale = min(SIZE / h, SIZE / w)
|
||||
image = image.resize((max(int(w*scale),SIZE),max(int(h*scale),SIZE)), Image.Resampling.BICUBIC)
|
||||
w, h = image.size
|
||||
if w > SIZE:
|
||||
left = (w - SIZE) // 2
|
||||
image = image.crop((left, left+SIZE, 0, SIZE))
|
||||
elif h > SIZE:
|
||||
top = (h - SIZE) // 2
|
||||
image = image.crop((0, SIZE, top, top+SIZE))
|
||||
|
||||
x = Tensor(np.array(image.convert('RGB')), device=self.std.device)
|
||||
x = x.permute(2, 0, 1).cast(dtypes.float32) / 255.0
|
||||
return (x - self.mean) / self.std
|
||||
|
||||
def encode_tokens(self, tokens:Tensor) -> Tensor:
|
||||
x = self.token_embedding(tokens)
|
||||
x = x + self.positional_embedding
|
||||
x = self.transformer(x, attn_mask=self.attn_mask)
|
||||
x = self.ln_final(x)
|
||||
x = x[Tensor.arange(x.shape[0], device=x.device), tokens.argmax(axis=-1)]
|
||||
x = x @ self.text_projection
|
||||
return x
|
||||
|
||||
def get_clip_score(self, tokens:Tensor, image:Tensor) -> Tensor:
|
||||
image_features: Tensor = self.visual(image)
|
||||
image_features /= image_features.square().sum(-1, keepdim=True).sqrt() # Frobenius Norm
|
||||
|
||||
text_features = self.encode_tokens(tokens)
|
||||
text_features /= text_features.square().sum(-1, keepdim=True).sqrt() # Frobenius Norm
|
||||
|
||||
return (image_features * text_features).sum(axis=-1)
|
||||
64
tinygrad_repo/extra/models/convnext.py
Normal file
64
tinygrad_repo/extra/models/convnext.py
Normal file
@@ -0,0 +1,64 @@
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.nn import Conv2d, LayerNorm, LayerNorm2d, Linear
|
||||
from tinygrad.helpers import fetch, get_child
|
||||
|
||||
class Block:
|
||||
def __init__(self, dim):
|
||||
self.dwconv = Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim)
|
||||
self.norm = LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = Linear(dim, 4 * dim)
|
||||
self.pwconv2 = Linear(4 * dim, dim)
|
||||
self.gamma = Tensor.ones(dim)
|
||||
|
||||
def __call__(self, x:Tensor):
|
||||
return x + x.sequential([
|
||||
self.dwconv, lambda x: x.permute(0, 2, 3, 1), self.norm,
|
||||
self.pwconv1, Tensor.gelu, self.pwconv2, lambda x: (self.gamma * x).permute(0, 3, 1, 2)
|
||||
])
|
||||
|
||||
class ConvNeXt:
|
||||
def __init__(self, in_chans=3, num_classes=1000, depths=[3, 3, 9, 3], dims=[96, 192, 384, 768]):
|
||||
self.downsample_layers = [
|
||||
[Conv2d(in_chans, dims[0], kernel_size=4, stride=4), LayerNorm2d(dims[0], eps=1e-6)],
|
||||
*[[LayerNorm2d(dims[i], eps=1e-6), Conv2d(dims[i], dims[i+1], kernel_size=2, stride=2)] for i in range(len(dims)-1)]
|
||||
]
|
||||
self.stages = [[Block(dims[i]) for _ in range(depths[i])] for i in range(len(dims))]
|
||||
self.norm = LayerNorm(dims[-1])
|
||||
self.head = Linear(dims[-1], num_classes)
|
||||
|
||||
def __call__(self, x:Tensor):
|
||||
for downsample, stage in zip(self.downsample_layers, self.stages):
|
||||
x = x.sequential(downsample).sequential(stage)
|
||||
return x.mean([-2, -1]).sequential([self.norm, self.head])
|
||||
|
||||
# *** model definition is done ***
|
||||
|
||||
versions = {
|
||||
"tiny": {"depths": [3, 3, 9, 3], "dims": [96, 192, 384, 768]},
|
||||
"small": {"depths": [3, 3, 27, 3], "dims": [96, 192, 384, 768]},
|
||||
"base": {"depths": [3, 3, 9, 3], "dims": [128, 256, 512, 1024]},
|
||||
"large": {"depths": [3, 3, 27, 3], "dims": [192, 384, 768, 1536]},
|
||||
"xlarge": {"depths": [3, 3, 27, 3], "dims": [256, 512, 1024, 2048]}
|
||||
}
|
||||
|
||||
def get_model(version, load_weights=False):
|
||||
model = ConvNeXt(**versions[version])
|
||||
if load_weights:
|
||||
from tinygrad.nn.state import torch_load
|
||||
weights = torch_load(fetch(f'https://dl.fbaipublicfiles.com/convnext/convnext_{version}_1k_224_ema.pth'))['model']
|
||||
for k,v in weights.items():
|
||||
mv = get_child(model, k)
|
||||
mv.assign(v.reshape(mv.shape).to(mv.device)).realize()
|
||||
return model
|
||||
|
||||
if __name__ == "__main__":
|
||||
model = get_model("tiny", True)
|
||||
|
||||
# load image
|
||||
from test.models.test_efficientnet import chicken_img, preprocess, _LABELS
|
||||
img = Tensor(preprocess(chicken_img))
|
||||
|
||||
Tensor.training = False
|
||||
|
||||
out = model(img).numpy()
|
||||
print(_LABELS[out.argmax()])
|
||||
164
tinygrad_repo/extra/models/efficientnet.py
Normal file
164
tinygrad_repo/extra/models/efficientnet.py
Normal file
@@ -0,0 +1,164 @@
|
||||
import math
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.nn import BatchNorm2d
|
||||
from tinygrad.helpers import get_child, fetch
|
||||
from tinygrad.nn.state import torch_load
|
||||
|
||||
class MBConvBlock:
|
||||
def __init__(self, kernel_size, strides, expand_ratio, input_filters, output_filters, se_ratio, has_se, track_running_stats=True):
|
||||
oup = expand_ratio * input_filters
|
||||
if expand_ratio != 1:
|
||||
self._expand_conv = Tensor.glorot_uniform(oup, input_filters, 1, 1)
|
||||
self._bn0 = BatchNorm2d(oup, track_running_stats=track_running_stats)
|
||||
else:
|
||||
self._expand_conv = None
|
||||
|
||||
self.strides = strides
|
||||
if strides == (2,2):
|
||||
self.pad = [(kernel_size-1)//2-1, (kernel_size-1)//2]*2
|
||||
else:
|
||||
self.pad = [(kernel_size-1)//2]*4
|
||||
|
||||
self._depthwise_conv = Tensor.glorot_uniform(oup, 1, kernel_size, kernel_size)
|
||||
self._bn1 = BatchNorm2d(oup, track_running_stats=track_running_stats)
|
||||
|
||||
self.has_se = has_se
|
||||
if self.has_se:
|
||||
num_squeezed_channels = max(1, int(input_filters * se_ratio))
|
||||
self._se_reduce = Tensor.glorot_uniform(num_squeezed_channels, oup, 1, 1)
|
||||
self._se_reduce_bias = Tensor.zeros(num_squeezed_channels)
|
||||
self._se_expand = Tensor.glorot_uniform(oup, num_squeezed_channels, 1, 1)
|
||||
self._se_expand_bias = Tensor.zeros(oup)
|
||||
|
||||
self._project_conv = Tensor.glorot_uniform(output_filters, oup, 1, 1)
|
||||
self._bn2 = BatchNorm2d(output_filters, track_running_stats=track_running_stats)
|
||||
|
||||
def __call__(self, inputs):
|
||||
x = inputs
|
||||
if self._expand_conv is not None:
|
||||
x = self._bn0(x.conv2d(self._expand_conv)).swish()
|
||||
x = x.conv2d(self._depthwise_conv, padding=self.pad, stride=self.strides, groups=self._depthwise_conv.shape[0])
|
||||
x = self._bn1(x).swish()
|
||||
|
||||
if self.has_se:
|
||||
x_squeezed = x.avg_pool2d(kernel_size=x.shape[2:4])
|
||||
x_squeezed = x_squeezed.conv2d(self._se_reduce, self._se_reduce_bias).swish()
|
||||
x_squeezed = x_squeezed.conv2d(self._se_expand, self._se_expand_bias)
|
||||
x = x.mul(x_squeezed.sigmoid())
|
||||
|
||||
x = self._bn2(x.conv2d(self._project_conv))
|
||||
if x.shape == inputs.shape:
|
||||
x = x.add(inputs)
|
||||
return x
|
||||
|
||||
class EfficientNet:
|
||||
def __init__(self, number=0, classes=1000, has_se=True, track_running_stats=True, input_channels=3, has_fc_output=True):
|
||||
self.number = number
|
||||
global_params = [
|
||||
# width, depth
|
||||
(1.0, 1.0), # b0
|
||||
(1.0, 1.1), # b1
|
||||
(1.1, 1.2), # b2
|
||||
(1.2, 1.4), # b3
|
||||
(1.4, 1.8), # b4
|
||||
(1.6, 2.2), # b5
|
||||
(1.8, 2.6), # b6
|
||||
(2.0, 3.1), # b7
|
||||
(2.2, 3.6), # b8
|
||||
(4.3, 5.3), # l2
|
||||
][max(number,0)]
|
||||
|
||||
def round_filters(filters):
|
||||
multiplier = global_params[0]
|
||||
divisor = 8
|
||||
filters *= multiplier
|
||||
new_filters = max(divisor, int(filters + divisor / 2) // divisor * divisor)
|
||||
if new_filters < 0.9 * filters: # prevent rounding by more than 10%
|
||||
new_filters += divisor
|
||||
return int(new_filters)
|
||||
|
||||
def round_repeats(repeats):
|
||||
return int(math.ceil(global_params[1] * repeats))
|
||||
|
||||
out_channels = round_filters(32)
|
||||
self._conv_stem = Tensor.glorot_uniform(out_channels, input_channels, 3, 3)
|
||||
self._bn0 = BatchNorm2d(out_channels, track_running_stats=track_running_stats)
|
||||
blocks_args = [
|
||||
[1, 3, (1,1), 1, 32, 16, 0.25],
|
||||
[2, 3, (2,2), 6, 16, 24, 0.25],
|
||||
[2, 5, (2,2), 6, 24, 40, 0.25],
|
||||
[3, 3, (2,2), 6, 40, 80, 0.25],
|
||||
[3, 5, (1,1), 6, 80, 112, 0.25],
|
||||
[4, 5, (2,2), 6, 112, 192, 0.25],
|
||||
[1, 3, (1,1), 6, 192, 320, 0.25],
|
||||
]
|
||||
|
||||
if self.number == -1:
|
||||
blocks_args = [
|
||||
[1, 3, (2,2), 1, 32, 40, 0.25],
|
||||
[1, 3, (2,2), 1, 40, 80, 0.25],
|
||||
[1, 3, (2,2), 1, 80, 192, 0.25],
|
||||
[1, 3, (2,2), 1, 192, 320, 0.25],
|
||||
]
|
||||
elif self.number == -2:
|
||||
blocks_args = [
|
||||
[1, 9, (8,8), 1, 32, 320, 0.25],
|
||||
]
|
||||
|
||||
self._blocks = []
|
||||
for num_repeats, kernel_size, strides, expand_ratio, input_filters, output_filters, se_ratio in blocks_args:
|
||||
input_filters, output_filters = round_filters(input_filters), round_filters(output_filters)
|
||||
for n in range(round_repeats(num_repeats)):
|
||||
self._blocks.append(MBConvBlock(kernel_size, strides, expand_ratio, input_filters, output_filters, se_ratio, has_se=has_se, track_running_stats=track_running_stats))
|
||||
input_filters = output_filters
|
||||
strides = (1,1)
|
||||
|
||||
in_channels = round_filters(320)
|
||||
out_channels = round_filters(1280)
|
||||
self._conv_head = Tensor.glorot_uniform(out_channels, in_channels, 1, 1)
|
||||
self._bn1 = BatchNorm2d(out_channels, track_running_stats=track_running_stats)
|
||||
if has_fc_output:
|
||||
self._fc = Tensor.glorot_uniform(out_channels, classes)
|
||||
self._fc_bias = Tensor.zeros(classes)
|
||||
else:
|
||||
self._fc = None
|
||||
|
||||
def forward(self, x):
|
||||
x = self._bn0(x.conv2d(self._conv_stem, padding=(0,1,0,1), stride=2)).swish()
|
||||
x = x.sequential(self._blocks)
|
||||
x = self._bn1(x.conv2d(self._conv_head)).swish()
|
||||
x = x.avg_pool2d(kernel_size=x.shape[2:4])
|
||||
x = x.reshape(shape=(-1, x.shape[1]))
|
||||
return x.linear(self._fc, self._fc_bias) if self._fc is not None else x
|
||||
|
||||
def load_from_pretrained(self):
|
||||
model_urls = {
|
||||
0: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b0-355c32eb.pth",
|
||||
1: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b1-f1951068.pth",
|
||||
2: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b2-8bb594d6.pth",
|
||||
3: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b3-5fb5a3c3.pth",
|
||||
4: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b4-6ed6700e.pth",
|
||||
5: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b5-b6417697.pth",
|
||||
6: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b6-c76e70fd.pth",
|
||||
7: "https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b7-dcc49843.pth"
|
||||
}
|
||||
|
||||
b0 = torch_load(fetch(model_urls[self.number]))
|
||||
for k,v in b0.items():
|
||||
if k.endswith("num_batches_tracked"): continue
|
||||
for cat in ['_conv_head', '_conv_stem', '_depthwise_conv', '_expand_conv', '_fc', '_project_conv', '_se_reduce', '_se_expand']:
|
||||
if cat in k:
|
||||
k = k.replace('.bias', '_bias')
|
||||
k = k.replace('.weight', '')
|
||||
|
||||
#print(k, v.shape)
|
||||
mv:Tensor = get_child(self, k)
|
||||
vnp = v #.astype(np.float32)
|
||||
vnp = vnp if k != '_fc' else vnp.T
|
||||
#vnp = vnp if vnp.shape != () else np.array([vnp])
|
||||
|
||||
if mv.shape == vnp.shape:
|
||||
mv.replace(vnp.to(mv.device))
|
||||
else:
|
||||
print("MISMATCH SHAPE IN %s, %r %r" % (k, mv.shape, vnp.shape))
|
||||
|
||||
344
tinygrad_repo/extra/models/inception.py
Normal file
344
tinygrad_repo/extra/models/inception.py
Normal file
@@ -0,0 +1,344 @@
|
||||
from tinygrad import Tensor
|
||||
from tinygrad.nn import Conv2d, BatchNorm2d, Linear
|
||||
from tinygrad.nn.state import load_state_dict, torch_load
|
||||
from tinygrad.helpers import fetch
|
||||
|
||||
from typing import Optional, Dict
|
||||
import numpy as np
|
||||
from scipy import linalg
|
||||
|
||||
# Base Inception Model
|
||||
|
||||
class BasicConv2d:
|
||||
def __init__(self, in_ch:int, out_ch:int, **kwargs):
|
||||
self.conv = Conv2d(in_ch, out_ch, bias=False, **kwargs)
|
||||
self.bn = BatchNorm2d(out_ch, eps=0.001)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return x.sequential([self.conv, self.bn, Tensor.relu])
|
||||
|
||||
class InceptionA:
|
||||
def __init__(self, in_ch:int, pool_feat:int):
|
||||
self.branch1x1 = BasicConv2d(in_ch, 64, kernel_size=1)
|
||||
|
||||
self.branch5x5_1 = BasicConv2d(in_ch, 48, kernel_size=1)
|
||||
self.branch5x5_2 = BasicConv2d(48, 64, kernel_size=5, padding=2)
|
||||
|
||||
self.branch3x3dbl_1 = BasicConv2d(in_ch, 64, kernel_size=1)
|
||||
self.branch3x3dbl_2 = BasicConv2d(64, 96, kernel_size=(3,3), padding=1)
|
||||
self.branch3x3dbl_3 = BasicConv2d(96, 96, kernel_size=(3,3), padding=1)
|
||||
|
||||
self.branch_pool = BasicConv2d(in_ch, pool_feat, kernel_size=1)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
x.sequential([self.branch5x5_1, self.branch5x5_2]),
|
||||
x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2, self.branch3x3dbl_3]),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1)),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class InceptionB:
|
||||
def __init__(self, in_ch:int):
|
||||
self.branch3x3 = BasicConv2d(in_ch, 384, kernel_size=(3,3), stride=2)
|
||||
|
||||
self.branch3x3dbl_1 = BasicConv2d(in_ch, 64, kernel_size=1)
|
||||
self.branch3x3dbl_2 = BasicConv2d(64, 96, kernel_size=(3,3), padding=1)
|
||||
self.branch3x3dbl_3 = BasicConv2d(96, 96, kernel_size=(3,3), stride=2)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
self.branch3x3(x),
|
||||
x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2, self.branch3x3dbl_3]),
|
||||
x.max_pool2d(kernel_size=(3,3), stride=2, dilation=1),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class InceptionC:
|
||||
def __init__(self, in_ch, ch_7x7):
|
||||
self.branch1x1 = BasicConv2d(in_ch, 192, kernel_size=1)
|
||||
|
||||
self.branch7x7_1 = BasicConv2d(in_ch, ch_7x7, kernel_size=1)
|
||||
self.branch7x7_2 = BasicConv2d(ch_7x7, ch_7x7, kernel_size=(1, 7), padding=(0, 3))
|
||||
self.branch7x7_3 = BasicConv2d(ch_7x7, 192, kernel_size=(7, 1), padding=(3, 0))
|
||||
|
||||
self.branch7x7dbl_1 = BasicConv2d(in_ch, ch_7x7, kernel_size=1)
|
||||
self.branch7x7dbl_2 = BasicConv2d(ch_7x7, ch_7x7, kernel_size=(7, 1), padding=(3, 0))
|
||||
self.branch7x7dbl_3 = BasicConv2d(ch_7x7, ch_7x7, kernel_size=(1, 7), padding=(0, 3))
|
||||
self.branch7x7dbl_4 = BasicConv2d(ch_7x7, ch_7x7, kernel_size=(7, 1), padding=(3, 0))
|
||||
self.branch7x7dbl_5 = BasicConv2d(ch_7x7, 192, kernel_size=(1, 7), padding=(0, 3))
|
||||
|
||||
self.branch_pool = BasicConv2d(in_ch, 192, kernel_size=1)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
x.sequential([self.branch7x7_1, self.branch7x7_2, self.branch7x7_3]),
|
||||
x.sequential([self.branch7x7dbl_1, self.branch7x7dbl_2, self.branch7x7dbl_3, self.branch7x7dbl_4, self.branch7x7dbl_5]),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1)),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class InceptionD:
|
||||
def __init__(self, in_ch:int):
|
||||
self.branch3x3_1 = BasicConv2d(in_ch, 192, kernel_size=1)
|
||||
self.branch3x3_2 = BasicConv2d(192, 320, kernel_size=(3,3), stride=2)
|
||||
|
||||
self.branch7x7x3_1 = BasicConv2d(in_ch, 192, kernel_size=1)
|
||||
self.branch7x7x3_2 = BasicConv2d(192, 192, kernel_size=(1, 7), padding=(0, 3))
|
||||
self.branch7x7x3_3 = BasicConv2d(192, 192, kernel_size=(7, 1), padding=(3, 0))
|
||||
self.branch7x7x3_4 = BasicConv2d(192, 192, kernel_size=(3,3), stride=2)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
x.sequential([self.branch3x3_1, self.branch3x3_2]),
|
||||
x.sequential([self.branch7x7x3_1, self.branch7x7x3_2, self.branch7x7x3_3, self.branch7x7x3_4]),
|
||||
x.max_pool2d(kernel_size=(3,3), stride=2, dilation=1),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class InceptionE:
|
||||
def __init__(self, in_ch:int):
|
||||
self.branch1x1 = BasicConv2d(in_ch, 320, kernel_size=1)
|
||||
|
||||
self.branch3x3_1 = BasicConv2d(in_ch, 384, kernel_size=1)
|
||||
self.branch3x3_2a = BasicConv2d(384, 384, kernel_size=(1, 3), padding=(0, 1))
|
||||
self.branch3x3_2b = BasicConv2d(384, 384, kernel_size=(3, 1), padding=(1, 0))
|
||||
|
||||
self.branch3x3dbl_1 = BasicConv2d(in_ch, 448, kernel_size=1)
|
||||
self.branch3x3dbl_2 = BasicConv2d(448, 384, kernel_size=(3,3), padding=1)
|
||||
self.branch3x3dbl_3a = BasicConv2d(384, 384, kernel_size=(1, 3), padding=(0, 1))
|
||||
self.branch3x3dbl_3b = BasicConv2d(384, 384, kernel_size=(3, 1), padding=(1, 0))
|
||||
|
||||
self.branch_pool = BasicConv2d(in_ch, 192, kernel_size=1)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
branch3x3 = self.branch3x3_1(x)
|
||||
branch3x3dbl = x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2])
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
Tensor.cat(self.branch3x3_2a(branch3x3), self.branch3x3_2b(branch3x3), dim=1),
|
||||
Tensor.cat(self.branch3x3dbl_3a(branch3x3dbl), self.branch3x3dbl_3b(branch3x3dbl), dim=1),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1)),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class InceptionAux:
|
||||
def __init__(self, in_ch:int, num_classes:int):
|
||||
self.conv0 = BasicConv2d(in_ch, 128, kernel_size=1)
|
||||
self.conv1 = BasicConv2d(128, 768, kernel_size=5)
|
||||
self.fc = Linear(768, num_classes)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
x = x.avg_pool2d(kernel_size=5, stride=3, padding=1).sequential([self.conv0, self.conv1])
|
||||
x = x.avg_pool2d(kernel_size=1, padding=1).reshape(x.shape[0],-1)
|
||||
return self.fc(x)
|
||||
|
||||
class Inception3:
|
||||
def __init__(self, num_classes:int=1008, cls_map:Optional[Dict]=None):
|
||||
def get_cls(key1:str, key2:str, default):
|
||||
return default if cls_map is None else cls_map.get(key1, cls_map.get(key2, default))
|
||||
|
||||
self.transform_input = False
|
||||
self.Conv2d_1a_3x3 = BasicConv2d(3, 32, kernel_size=(3,3), stride=2)
|
||||
self.Conv2d_2a_3x3 = BasicConv2d(32, 32, kernel_size=(3,3))
|
||||
self.Conv2d_2b_3x3 = BasicConv2d(32, 64, kernel_size=(3,3), padding=1)
|
||||
self.maxpool1 = lambda x: Tensor.max_pool2d(x, kernel_size=(3,3), stride=2, padding=1)
|
||||
self.Conv2d_3b_1x1 = BasicConv2d(64, 80, kernel_size=1)
|
||||
self.Conv2d_4a_3x3 = BasicConv2d(80, 192, kernel_size=(3,3))
|
||||
self.maxpool2 = lambda x: Tensor.max_pool2d(x, kernel_size=(3,3), stride=2, padding=1)
|
||||
self.Mixed_5b = get_cls("A1","A",InceptionA)(192, pool_feat=32)
|
||||
self.Mixed_5c = get_cls("A2","A",InceptionA)(256, pool_feat=64)
|
||||
self.Mixed_5d = get_cls("A3","A",InceptionA)(288, pool_feat=64)
|
||||
self.Mixed_6a = get_cls("B1","B",InceptionB)(288)
|
||||
self.Mixed_6b = get_cls("C1","C",InceptionC)(768, ch_7x7=128)
|
||||
self.Mixed_6c = get_cls("C2","C",InceptionC)(768, ch_7x7=160)
|
||||
self.Mixed_6d = get_cls("C3","C",InceptionC)(768, ch_7x7=160)
|
||||
self.Mixed_6e = get_cls("C4","C",InceptionC)(768, ch_7x7=192)
|
||||
self.Mixed_7a = get_cls("D1","D",InceptionD)(768)
|
||||
self.Mixed_7b = get_cls("E1","E",InceptionE)(1280)
|
||||
self.Mixed_7c = get_cls("E2","E",InceptionE)(2048)
|
||||
self.avgpool = lambda x: Tensor.avg_pool2d(x, kernel_size=(8,8), padding=1)
|
||||
self.fc = Linear(2048, num_classes)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return x.sequential([
|
||||
self.Conv2d_1a_3x3,
|
||||
self.Conv2d_2a_3x3,
|
||||
self.Conv2d_2b_3x3,
|
||||
self.maxpool1,
|
||||
|
||||
self.Conv2d_3b_1x1,
|
||||
self.Conv2d_4a_3x3,
|
||||
self.maxpool2,
|
||||
|
||||
self.Mixed_5b,
|
||||
self.Mixed_5c,
|
||||
self.Mixed_5d,
|
||||
self.Mixed_6a,
|
||||
self.Mixed_6b,
|
||||
self.Mixed_6c,
|
||||
self.Mixed_6d,
|
||||
self.Mixed_6e,
|
||||
|
||||
self.Mixed_7a,
|
||||
self.Mixed_7b,
|
||||
self.Mixed_7c,
|
||||
self.avgpool,
|
||||
|
||||
lambda y: y.reshape(x.shape[0],-1),
|
||||
self.fc,
|
||||
])
|
||||
|
||||
|
||||
# FID Inception Variation
|
||||
|
||||
class FidInceptionA(InceptionA):
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
x.sequential([self.branch5x5_1, self.branch5x5_2]),
|
||||
x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2, self.branch3x3dbl_3]),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1, count_include_pad=False))
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class FidInceptionC(InceptionC):
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
x.sequential([self.branch7x7_1, self.branch7x7_2, self.branch7x7_3]),
|
||||
x.sequential([self.branch7x7dbl_1, self.branch7x7dbl_2, self.branch7x7dbl_3, self.branch7x7dbl_4, self.branch7x7dbl_5]),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1, count_include_pad=False))
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class FidInceptionE1(InceptionE):
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
branch3x3 = self.branch3x3_1(x)
|
||||
branch3x3dbl = x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2])
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
Tensor.cat(self.branch3x3_2a(branch3x3), self.branch3x3_2b(branch3x3), dim=1),
|
||||
Tensor.cat(self.branch3x3dbl_3a(branch3x3dbl), self.branch3x3dbl_3b(branch3x3dbl), dim=1),
|
||||
self.branch_pool(x.avg_pool2d(kernel_size=(3,3), stride=1, padding=1, count_include_pad=False)),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class FidInceptionE2(InceptionE):
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
branch3x3 = self.branch3x3_1(x)
|
||||
branch3x3dbl = x.sequential([self.branch3x3dbl_1, self.branch3x3dbl_2])
|
||||
outputs = [
|
||||
self.branch1x1(x),
|
||||
Tensor.cat(self.branch3x3_2a(branch3x3), self.branch3x3_2b(branch3x3), dim=1),
|
||||
Tensor.cat(self.branch3x3dbl_3a(branch3x3dbl), self.branch3x3dbl_3b(branch3x3dbl), dim=1),
|
||||
self.branch_pool(x.max_pool2d(kernel_size=(3,3), stride=1, padding=1)),
|
||||
]
|
||||
return Tensor.cat(*outputs, dim=1)
|
||||
|
||||
class FidInceptionV3:
|
||||
m1: Optional[np.ndarray] = None
|
||||
s1: Optional[np.ndarray] = None
|
||||
|
||||
def __init__(self):
|
||||
inception = Inception3(cls_map={
|
||||
"A": FidInceptionA,
|
||||
"C": FidInceptionC,
|
||||
"E1": FidInceptionE1,
|
||||
"E2": FidInceptionE2,
|
||||
})
|
||||
|
||||
self.Conv2d_1a_3x3 = inception.Conv2d_1a_3x3
|
||||
self.Conv2d_2a_3x3 = inception.Conv2d_2a_3x3
|
||||
self.Conv2d_2b_3x3 = inception.Conv2d_2b_3x3
|
||||
|
||||
self.Conv2d_3b_1x1 = inception.Conv2d_3b_1x1
|
||||
self.Conv2d_4a_3x3 = inception.Conv2d_4a_3x3
|
||||
|
||||
self.Mixed_5b = inception.Mixed_5b
|
||||
self.Mixed_5c = inception.Mixed_5c
|
||||
self.Mixed_5d = inception.Mixed_5d
|
||||
self.Mixed_6a = inception.Mixed_6a
|
||||
self.Mixed_6b = inception.Mixed_6b
|
||||
self.Mixed_6c = inception.Mixed_6c
|
||||
self.Mixed_6d = inception.Mixed_6d
|
||||
self.Mixed_6e = inception.Mixed_6e
|
||||
|
||||
self.Mixed_7a = inception.Mixed_7a
|
||||
self.Mixed_7b = inception.Mixed_7b
|
||||
self.Mixed_7c = inception.Mixed_7c
|
||||
|
||||
def load_from_pretrained(self, path=None):
|
||||
if path is None:
|
||||
path = fetch("https://github.com/mseitzer/pytorch-fid/releases/download/fid_weights/pt_inception-2015-12-05-6726825d.pth", "pt_inception-2015-12-05-6726825d.pth")
|
||||
state_dict = torch_load(str(path))
|
||||
for k,v in state_dict.items():
|
||||
if k.endswith(".num_batches_tracked"):
|
||||
state_dict[k] = v.reshape(1)
|
||||
load_state_dict(self, state_dict)
|
||||
return self
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
x = x.interpolate((299,299), mode="linear")
|
||||
x = (x * 2) - 1
|
||||
x = x.sequential([
|
||||
self.Conv2d_1a_3x3,
|
||||
self.Conv2d_2a_3x3,
|
||||
self.Conv2d_2b_3x3,
|
||||
lambda x: Tensor.max_pool2d(x, kernel_size=(3,3), stride=2, dilation=1),
|
||||
|
||||
self.Conv2d_3b_1x1,
|
||||
self.Conv2d_4a_3x3,
|
||||
lambda x: Tensor.max_pool2d(x, kernel_size=(3,3), stride=2, dilation=1),
|
||||
|
||||
self.Mixed_5b,
|
||||
self.Mixed_5c,
|
||||
self.Mixed_5d,
|
||||
self.Mixed_6a,
|
||||
self.Mixed_6b,
|
||||
self.Mixed_6c,
|
||||
self.Mixed_6d,
|
||||
self.Mixed_6e,
|
||||
|
||||
self.Mixed_7a,
|
||||
self.Mixed_7b,
|
||||
self.Mixed_7c,
|
||||
lambda x: Tensor.avg_pool2d(x, kernel_size=(8,8)),
|
||||
])
|
||||
return x
|
||||
|
||||
def compute_score(self, inception_activations:Tensor, val_stats_path:str) -> float:
|
||||
if self.m1 is None and self.s1 is None:
|
||||
with np.load(val_stats_path) as f:
|
||||
self.m1, self.s1 = f['mu'][:], f['sigma'][:]
|
||||
assert self.m1 is not None and self.s1 is not None
|
||||
|
||||
m2 = inception_activations.mean(axis=0).numpy()
|
||||
s2 = np.cov(inception_activations.numpy(), rowvar=False)
|
||||
|
||||
return calculate_frechet_distance(self.m1, self.s1, m2, s2)
|
||||
|
||||
def calculate_frechet_distance(mu1:np.ndarray, sigma1:np.ndarray, mu2:np.ndarray, sigma2:np.ndarray, eps:float=1e-6) -> float:
|
||||
mu1 = np.atleast_1d(mu1)
|
||||
mu2 = np.atleast_1d(mu2)
|
||||
sigma1 = np.atleast_2d(sigma1)
|
||||
sigma2 = np.atleast_2d(sigma2)
|
||||
assert mu1.shape == mu2.shape and sigma1.shape == sigma2.shape
|
||||
|
||||
diff = mu1 - mu2
|
||||
covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
|
||||
if not np.isfinite(covmean).all():
|
||||
offset = np.eye(sigma1.shape[0]) * eps
|
||||
covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
|
||||
|
||||
if np.iscomplexobj(covmean):
|
||||
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
|
||||
m = np.max(np.abs(covmean.imag))
|
||||
raise ValueError(f"Imaginary component {m}")
|
||||
covmean = covmean.real
|
||||
|
||||
tr_covmean = np.trace(covmean)
|
||||
|
||||
return diff.dot(diff) + np.trace(sigma1) + np.trace(sigma2) - 2*tr_covmean
|
||||
286
tinygrad_repo/extra/models/llama.py
Normal file
286
tinygrad_repo/extra/models/llama.py
Normal file
@@ -0,0 +1,286 @@
|
||||
from typing import Union, Optional, Any
|
||||
import collections, math
|
||||
from tinygrad import Tensor, Variable, TinyJit, dtypes, nn, Device
|
||||
from tinygrad.helpers import getenv, DEBUG
|
||||
|
||||
# https://github.com/facebookresearch/llama/blob/1076b9c51c77ad06e9d7ba8a4c6df775741732bd/llama/model.py#L47
|
||||
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0) -> Tensor:
|
||||
freqs = 1.0 / (theta ** (Tensor.arange(0, dim, 2)[:(dim // 2)] / dim))
|
||||
freqs = Tensor.arange(end).unsqueeze(dim=1) * freqs.unsqueeze(dim=0)
|
||||
return Tensor.stack(freqs.cos(), freqs.sin(), dim=-1).reshape(1, end, 1, dim//2, 2)
|
||||
|
||||
# matches meta, non hugging face weights
|
||||
# (a+i*b) * (c+i*d) = (ac-bd) + i*(ad+bc)
|
||||
def complex_mult(A, c, d):
|
||||
a,b = A[..., 0:1], A[..., 1:2]
|
||||
ro = a*c - b*d
|
||||
co = a*d + b*c
|
||||
return ro.cat(co, dim=-1)
|
||||
|
||||
def apply_rotary_emb(xq:Tensor, xk:Tensor, freqs_cis:Tensor) -> tuple[Tensor, Tensor]:
|
||||
assert freqs_cis.shape[1] == xq.shape[1] == xk.shape[1], f"freqs_cis shape mismatch {freqs_cis.shape} xq:{xq.shape} xk:{xk.shape}"
|
||||
xq = xq.reshape(*xq.shape[0:-1], -1, 2)
|
||||
xk = xk.reshape(*xk.shape[0:-1], -1, 2)
|
||||
assert len(xq.shape) == len(xk.shape) == len(freqs_cis.shape) == 5
|
||||
c, d = freqs_cis[..., 0:1], freqs_cis[..., 1:2]
|
||||
xq_out = complex_mult(xq, c, d)
|
||||
xk_out = complex_mult(xk, c, d)
|
||||
return xq_out.flatten(3), xk_out.flatten(3)
|
||||
|
||||
def repeat_kv(x:Tensor, n_rep:int) -> Tensor:
|
||||
bs, seqlen, n_kv_heads, head_dim = x.shape
|
||||
if n_rep == 1: return x
|
||||
# NOTE: this is different from x.repeat((1, 1, n_rep, 1))
|
||||
return x.repeat((1, 1, 1, n_rep)).reshape(bs, seqlen, n_kv_heads * n_rep, head_dim)
|
||||
|
||||
class Attention:
|
||||
def __init__(self, dim, n_heads, n_kv_heads=None, max_context=0, linear=nn.Linear, qk_norm:float|None=None):
|
||||
self.n_heads = n_heads
|
||||
self.n_kv_heads = n_kv_heads if n_kv_heads is not None else n_heads # n_kv_heads != n_heads implies MQA [arxiv/2307.09288, A.2.1]
|
||||
self.head_dim = dim // n_heads
|
||||
self.n_rep = self.n_heads // self.n_kv_heads
|
||||
self.max_context = max_context
|
||||
|
||||
if getenv("WQKV"):
|
||||
self.wqkv = linear(dim, self.n_heads * self.head_dim + self.n_kv_heads * self.head_dim * 2, bias=False)
|
||||
else:
|
||||
self.wq = linear(dim, self.n_heads * self.head_dim, bias=False)
|
||||
self.wk = linear(dim, self.n_kv_heads * self.head_dim, bias=False)
|
||||
self.wv = linear(dim, self.n_kv_heads * self.head_dim, bias=False)
|
||||
|
||||
self.wo = linear(self.n_heads * self.head_dim, dim, bias=False)
|
||||
|
||||
self.q_norm = nn.RMSNorm(dim, qk_norm) if qk_norm is not None else None
|
||||
self.k_norm = nn.RMSNorm(dim, qk_norm) if qk_norm is not None else None
|
||||
|
||||
def __call__(self, x:Tensor, start_pos:Union[Variable,int], freqs_cis:Tensor, mask:Optional[Tensor]=None) -> Tensor:
|
||||
if getenv("WQKV"):
|
||||
xqkv = self.wqkv(x)
|
||||
xqkv = xqkv.reshape(xqkv.shape[0], xqkv.shape[1], self.n_kv_heads, self.n_rep + 2, self.head_dim)
|
||||
xq = xqkv[:, :, :, :self.n_rep].reshape(xqkv.shape[0], xqkv.shape[1], -1)
|
||||
xk = xqkv[:, :, :, self.n_rep:self.n_rep+1].reshape(xqkv.shape[0], xqkv.shape[1], -1)
|
||||
xv = xqkv[:, :, :, self.n_rep+1:self.n_rep+2].reshape(xqkv.shape[0], xqkv.shape[1], -1)
|
||||
else:
|
||||
xq, xk, xv = self.wq(x), self.wk(x.contiguous_backward()), self.wv(x)
|
||||
|
||||
if self.q_norm is not None and self.k_norm is not None:
|
||||
xq = self.q_norm(xq)
|
||||
xk = self.k_norm(xk)
|
||||
|
||||
# cast_float_to_bf16 is expensive in reduction loops, break it out
|
||||
if x.dtype == dtypes.bfloat16: xq, xk = xq.contiguous_backward(), xk.contiguous_backward()
|
||||
|
||||
xq = xq.reshape(xq.shape[0], xq.shape[1], self.n_heads, self.head_dim)
|
||||
xk = xk.reshape(xk.shape[0], xk.shape[1], self.n_kv_heads, self.head_dim)
|
||||
xv = xv.reshape(xv.shape[0], xv.shape[1], self.n_kv_heads, self.head_dim)
|
||||
|
||||
xq, xk = apply_rotary_emb(xq, xk, freqs_cis)
|
||||
bsz, seqlen, _, _ = xq.shape
|
||||
|
||||
# create kv cache
|
||||
if self.max_context:
|
||||
if not hasattr(self, "cache_kv"):
|
||||
self.cache_kv = Tensor.zeros(2, bsz, self.max_context, self.n_kv_heads, self.head_dim, dtype=x.dtype).contiguous().realize()
|
||||
if isinstance(x.device, tuple):
|
||||
# TODO: instead of specifying how to shard, it can follow how xk and xv are being sharded
|
||||
self.cache_kv.shard_((x.device), axis=3 if getenv("SHARD_KVCACHE") else None).realize()
|
||||
|
||||
# update the cache
|
||||
assert xk.dtype == xv.dtype == self.cache_kv.dtype, f"{xk.dtype=}, {xv.dtype=}, {self.cache_kv.dtype=}"
|
||||
self.cache_kv[:, :, start_pos:start_pos+seqlen, :, :].assign(Tensor.stack(xk, xv)).realize()
|
||||
|
||||
keys = self.cache_kv[0, :, 0:start_pos+seqlen, :, :]
|
||||
values = self.cache_kv[1, :, 0:start_pos+seqlen, :, :]
|
||||
else:
|
||||
assert start_pos == 0
|
||||
keys, values = xk, xv
|
||||
|
||||
if self.max_context:
|
||||
keys, values = repeat_kv(keys, self.n_rep), repeat_kv(values, self.n_rep)
|
||||
xq, keys, values = xq.transpose(1, 2), keys.transpose(1, 2), values.transpose(1, 2)
|
||||
attn = xq.scaled_dot_product_attention(keys, values, mask).transpose(1, 2)
|
||||
else:
|
||||
xq, keys, values = xq.transpose(1, 2), keys.transpose(1, 2), values.transpose(1, 2)
|
||||
attn = xq.scaled_dot_product_attention(keys, values, is_causal=True, enable_gqa=True).transpose(1, 2)
|
||||
if getenv("STUB_ATTENTION"):
|
||||
from tinygrad.uop.ops import UOp, KernelInfo
|
||||
def fa_custom_forward(attn:UOp, q:UOp, k:UOp, v:UOp) -> UOp:
|
||||
return UOp.sink(arg=KernelInfo(name="fa_custom_forward"))
|
||||
def fa_custom_backward(out_q:UOp, out_k:UOp, out_v:UOp, grad:UOp, q:UOp, k:UOp, v:UOp) -> UOp:
|
||||
return UOp.sink(arg=KernelInfo(name="fa_custom_backward"))
|
||||
def fa_backward(grad:UOp, kernel:UOp) -> tuple[None, UOp, UOp, UOp]:
|
||||
grad_q = Tensor.empty_like(q:=Tensor(kernel.src[2]))
|
||||
grad_k = Tensor.empty_like(k:=Tensor(kernel.src[3]))
|
||||
grad_v = Tensor.empty_like(v:=Tensor(kernel.src[4]))
|
||||
ck = Tensor.custom_kernel(grad_q, grad_k, grad_v, Tensor(grad), q, k, v, fxn=fa_custom_backward)[:3]
|
||||
return (None, ck[0].uop, ck[1].uop, ck[2].uop)
|
||||
attn = Tensor.empty_like(attn).custom_kernel(xq, keys, values, fxn=fa_custom_forward, grad_fxn=fa_backward)[0]
|
||||
attn = attn.reshape(bsz, seqlen, -1)
|
||||
return self.wo(attn)
|
||||
|
||||
class FeedForward:
|
||||
def __init__(self, dim:int, hidden_dim:int, linear=nn.Linear):
|
||||
self.w1 = linear(dim, hidden_dim, bias=False)
|
||||
self.w2 = linear(hidden_dim, dim, bias=False)
|
||||
self.w3 = linear(dim, hidden_dim, bias=False) # the gate in Gated Linear Unit
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
w1 = self.w1(x).silu()
|
||||
w3 = self.w3(x.contiguous_backward()) # this fixes a strange fusion that makes tensor cores miss
|
||||
return self.w2(w1 * w3)
|
||||
|
||||
class TransformerBlock:
|
||||
def __init__(self, dim:int, hidden_dim:int, n_heads:int, n_kv_heads:int, norm_eps:float, max_context:int, linear=nn.Linear,
|
||||
feed_forward=FeedForward, qk_norm=None):
|
||||
self.attention = Attention(dim, n_heads, n_kv_heads, max_context, linear, qk_norm)
|
||||
self.feed_forward = feed_forward(dim, hidden_dim, linear)
|
||||
self.attention_norm = nn.RMSNorm(dim, norm_eps)
|
||||
self.ffn_norm = nn.RMSNorm(dim, norm_eps)
|
||||
|
||||
def __call__(self, x:Tensor, start_pos:Union[Variable,int], freqs_cis:Tensor, mask:Optional[Tensor]):
|
||||
h = x + self.attention(self.attention_norm(x), start_pos, freqs_cis, mask)
|
||||
return (h + self.feed_forward(self.ffn_norm(h))).contiguous().contiguous_backward()
|
||||
|
||||
# standard openai sampling
|
||||
def sample(logits: Tensor, temp: float, k: int, p: float, af: float, ap: float):
|
||||
assert logits.ndim == 1, "only works on 1d tensors"
|
||||
assert 0 <= p <= 1, "p must be between 0 and 1"
|
||||
assert 0 <= k <= logits.numel(), "k must be between 0 and numel"
|
||||
|
||||
# if temperature is very low just use argmax
|
||||
if temp < 1e-6: return logits.argmax()
|
||||
|
||||
logits = logits.to(Device.DEFAULT)
|
||||
|
||||
# alpha sampling
|
||||
if af or ap:
|
||||
if not hasattr(sample, "alpha_counter"):
|
||||
setattr(sample, "alpha_counter", Tensor.zeros_like(logits, dtype=dtypes.int32).contiguous())
|
||||
logits = logits - (sample.alpha_counter * af + (sample.alpha_counter > 0) * ap)
|
||||
|
||||
# replace NaNs with -inf
|
||||
logits = (logits != logits).where(-float("inf"), logits)
|
||||
|
||||
# softmax
|
||||
t = (logits / temp).softmax()
|
||||
|
||||
counter, counter2 = Tensor.arange(t.numel(), device=logits.device).contiguous(), Tensor.arange(t.numel() - 1, -1, -1, device=logits.device).contiguous()
|
||||
# top k
|
||||
if k:
|
||||
output, output_indices = Tensor.zeros(k, device=logits.device).contiguous(), Tensor.zeros(k, device=logits.device, dtype=dtypes.int32).contiguous()
|
||||
for i in range(k):
|
||||
t_argmax = (t.numel() - ((t == (t_max := t.max())) * counter2).max() - 1).cast(dtypes.default_int)
|
||||
output = output + t_max.unsqueeze(0).pad(((i, k - i - 1),))
|
||||
output_indices = output_indices + t_argmax.unsqueeze(0).pad(((i, k - i - 1),))
|
||||
t = (counter == t_argmax).where(0, t)
|
||||
|
||||
# approximate top p
|
||||
# because we are already limited to top k elements we can do top p "without sorting"
|
||||
output_cumsum = output[::-1].cumsum()[::-1] + t.sum()
|
||||
output = (output_cumsum >= (1 - p)) * output
|
||||
output_indices = (output_cumsum >= (1 - p)) * output_indices
|
||||
|
||||
# sample
|
||||
output_idx = output.multinomial()
|
||||
output_token = output_indices[output_idx]
|
||||
else:
|
||||
output_token = t.multinomial()
|
||||
|
||||
# increase alpha counter
|
||||
if af or ap:
|
||||
sample.alpha_counter = (counter == output_token).where(sample.alpha_counter + 1, sample.alpha_counter)
|
||||
|
||||
return output_token
|
||||
|
||||
class Transformer:
|
||||
def __init__(self, dim:int, hidden_dim:int, n_heads:int, n_layers:int, norm_eps:float, vocab_size, linear=nn.Linear, embedding=nn.Embedding,
|
||||
n_kv_heads=None, rope_theta=10000, max_context=1024, jit=True, feed_forward=FeedForward, qk_norm=None, disable_kv_cache=False):
|
||||
self.layers = [TransformerBlock(dim, hidden_dim, n_heads, n_kv_heads, norm_eps, 0 if disable_kv_cache else max_context,
|
||||
linear, feed_forward=feed_forward, qk_norm=qk_norm) for _ in range(n_layers)]
|
||||
self.norm = nn.RMSNorm(dim, norm_eps)
|
||||
self.tok_embeddings = embedding(vocab_size, dim)
|
||||
self.output = nn.Linear(dim, vocab_size, bias=False) if embedding == nn.Embedding else linear(dim, vocab_size, bias=False)
|
||||
self.max_context = max_context
|
||||
self.freqs_cis = precompute_freqs_cis(dim // n_heads, self.max_context * 2, rope_theta).contiguous().is_param_(False)
|
||||
self.forward_jit = TinyJit(self.forward) if jit else None
|
||||
|
||||
def forward(self, tokens:Tensor, start_pos:Union[Variable,int], temperature:float, top_k:int, top_p:float, alpha_f:float, alpha_p:float):
|
||||
_bsz, seqlen = tokens.shape
|
||||
h = self.tok_embeddings(tokens).contiguous()
|
||||
freqs_cis = self.freqs_cis.cast(h.dtype)[:, start_pos:start_pos+seqlen, :, :, :]
|
||||
|
||||
if self.max_context != 0 and seqlen > 1:
|
||||
mask = Tensor.full((1, 1, seqlen, start_pos+seqlen), float("-inf"), dtype=h.dtype, device=h.device).triu(start_pos+1)
|
||||
else: mask = None
|
||||
for layer in self.layers: h = layer(h, start_pos, freqs_cis, mask)
|
||||
logits = self.output(self.norm(h).contiguous().contiguous_backward()).contiguous_backward()
|
||||
if math.isnan(temperature): return logits
|
||||
|
||||
return sample(logits[:, -1, :].flatten(), temperature, top_k, top_p, alpha_f, alpha_p)
|
||||
|
||||
def __call__(self, tokens:Tensor, start_pos:int, temperature:float=0.0, top_k:int=0, top_p:float=0.8, alpha_f:float=0.0, alpha_p:float=0.0):
|
||||
# TODO: better way to handle the first call v.s. the rest?
|
||||
if tokens.shape[0:2] == (1,1) and self.forward_jit is not None and start_pos != 0:
|
||||
return self.forward_jit(tokens, Variable("start_pos", 1, self.max_context-1).bind(start_pos), temperature, top_k, top_p, alpha_f, alpha_p)
|
||||
return self.forward(tokens, start_pos, temperature, top_k, top_p, alpha_f, alpha_p)
|
||||
|
||||
# *** helpers ***
|
||||
|
||||
# TODO: n_kv_heads should support None
|
||||
def convert_from_huggingface(weights:dict[str, Tensor], n_layers: int, n_heads: int, n_kv_heads: int, permute_layers: bool = True):
|
||||
# huggingface stores Q and K permuted! it is mostly correct without this, but without it makes RoPE different, so it will diverge after 10+ toks.
|
||||
def permute(v: Tensor, n_heads: int):
|
||||
return v.reshape(n_heads, 2, v.shape[0] // n_heads // 2, v.shape[1] if len(v.shape) > 1 else 1).transpose(1, 2).reshape(*v.shape[:2])
|
||||
|
||||
keymap = {
|
||||
"model.embed_tokens.weight": "tok_embeddings.weight",
|
||||
**{f"model.layers.{l}.input_layernorm.weight": f"layers.{l}.attention_norm.weight" for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.self_attn.{x}_norm.weight": f"layers.{l}.attention.{x}_norm.weight" for x in ["q", "k"] for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.self_attn.{x}_proj.weight": f"layers.{l}.attention.w{x}.weight" for x in ["q", "k", "v", "o"] for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.self_attn.{x}_proj.bias": f"layers.{l}.attention.w{x}.bias" for x in ["q", "k", "v", "o"] for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.post_attention_layernorm.weight": f"layers.{l}.ffn_norm.weight" for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.mlp.{x}_proj.weight": f"layers.{l}.feed_forward.w{y}.weight" for x, y in {"gate": "1", "down": "2", "up": "3"}.items() for l in range(n_layers)},
|
||||
**{f"model.layers.{l}.mlp.gate.weight": f"layers.{l}.feed_forward.gate.weight" for l in range(n_layers)},
|
||||
"model.norm.weight": "norm.weight",
|
||||
"lm_head.weight": "output.weight",
|
||||
}
|
||||
sd = {}
|
||||
experts = collections.defaultdict(dict)
|
||||
for k, v in weights.items():
|
||||
if ".rotary_emb." in k: continue
|
||||
v = v.to(Device.DEFAULT)
|
||||
if "model.layers" in k:
|
||||
if ("q_proj" in k or "q_norm" in k) and permute_layers: v = permute(v, n_heads)
|
||||
elif ("k_proj" in k or "k_norm" in k) and permute_layers: v = permute(v, n_kv_heads)
|
||||
if '.mlp.experts.' in k:
|
||||
# support MoE models
|
||||
_, _, layer, _, _, expert, name, _ = k.split('.')
|
||||
experts[f'layers.{layer}.feed_forward.{name}'][int(expert)] = v
|
||||
continue
|
||||
sd[keymap[k]] = v
|
||||
for k,v in experts.items(): sd[k] = Tensor.stack(*[v[i] for i in range(len(v))])
|
||||
|
||||
# Handle tied embeddings (e.g., Llama 3.2 1B Instruct where lm_head shares weights with embed_tokens)
|
||||
if "output.weight" not in sd and "tok_embeddings.weight" in sd:
|
||||
sd["output.weight"] = sd["tok_embeddings.weight"]
|
||||
|
||||
return sd
|
||||
|
||||
def convert_from_gguf(weights:dict[str, Tensor], n_layers:int):
|
||||
keymap = {
|
||||
"token_embd.weight": "tok_embeddings.weight",
|
||||
**{f"blk.{l}.attn_norm.weight": f"layers.{l}.attention_norm.weight" for l in range(n_layers)},
|
||||
**{f"blk.{l}.attn_{x}.weight": f"layers.{l}.attention.w{x}.weight" for x in ["q", "k", "v"] for l in range(n_layers)},
|
||||
**{f"blk.{l}.attn_output.weight": f"layers.{l}.attention.wo.weight" for l in range(n_layers)},
|
||||
**{f"blk.{l}.ffn_norm.weight": f"layers.{l}.ffn_norm.weight" for l in range(n_layers)},
|
||||
**{f"blk.{l}.ffn_{x}.weight": f"layers.{l}.feed_forward.w{y}.weight" for x, y in {"gate": "1", "down": "2", "up": "3"}.items() for l in range(n_layers)},
|
||||
"output_norm.weight": "norm.weight",
|
||||
"rope_freqs.weight": "rope_freqs.weight",
|
||||
}
|
||||
sd = {keymap[k]: v for k,v in weights.items()}
|
||||
sd["output.weight"] = weights["token_embd.weight"]
|
||||
return sd
|
||||
|
||||
def fix_bf16(weights:dict[Any, Tensor]):
|
||||
# TODO: without casting to float16, 70B llama OOM on tinybox.
|
||||
return {k:v.cast(dtypes.float32).cast(dtypes.float16) if v.dtype == dtypes.bfloat16 else v for k,v in weights.items()}
|
||||
1214
tinygrad_repo/extra/models/mask_rcnn.py
Normal file
1214
tinygrad_repo/extra/models/mask_rcnn.py
Normal file
File diff suppressed because it is too large
Load Diff
170
tinygrad_repo/extra/models/resnet.py
Normal file
170
tinygrad_repo/extra/models/resnet.py
Normal file
@@ -0,0 +1,170 @@
|
||||
import tinygrad.nn as nn
|
||||
from tinygrad import Tensor, dtypes
|
||||
from tinygrad.nn.state import torch_load
|
||||
from tinygrad.helpers import fetch, get_child
|
||||
|
||||
# allow monkeypatching in layer implementations
|
||||
BatchNorm = nn.BatchNorm2d
|
||||
Conv2d = nn.Conv2d
|
||||
Linear = nn.Linear
|
||||
|
||||
|
||||
class BasicBlock:
|
||||
expansion = 1
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1, groups=1, base_width=64):
|
||||
assert groups == 1 and base_width == 64, "BasicBlock only supports groups=1 and base_width=64"
|
||||
self.conv1 = Conv2d(in_planes, planes, kernel_size=3, stride=stride, padding=1, bias=False)
|
||||
self.bn1 = BatchNorm(planes)
|
||||
self.conv2 = Conv2d(planes, planes, kernel_size=3, padding=1, stride=1, bias=False)
|
||||
self.bn2 = BatchNorm(planes)
|
||||
self.downsample = []
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.downsample = [
|
||||
Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False),
|
||||
BatchNorm(self.expansion*planes)
|
||||
]
|
||||
|
||||
def __call__(self, x):
|
||||
out = self.bn1(self.conv1(x)).relu()
|
||||
out = self.bn2(self.conv2(out))
|
||||
out = out + x.sequential(self.downsample)
|
||||
out = out.relu()
|
||||
return out
|
||||
|
||||
|
||||
class Bottleneck:
|
||||
# NOTE: stride_in_1x1=False, this is the v1.5 variant
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, in_planes, planes, stride=1, stride_in_1x1=False, groups=1, base_width=64):
|
||||
width = int(planes * (base_width / 64.0)) * groups
|
||||
# NOTE: the original implementation places stride at the first convolution (self.conv1), control with stride_in_1x1
|
||||
self.conv1 = Conv2d(in_planes, width, kernel_size=1, stride=stride if stride_in_1x1 else 1, bias=False)
|
||||
self.bn1 = BatchNorm(width)
|
||||
self.conv2 = Conv2d(width, width, kernel_size=3, padding=1, stride=1 if stride_in_1x1 else stride, groups=groups, bias=False)
|
||||
self.bn2 = BatchNorm(width)
|
||||
self.conv3 = Conv2d(width, self.expansion*planes, kernel_size=1, bias=False)
|
||||
self.bn3 = BatchNorm(self.expansion*planes)
|
||||
self.downsample = []
|
||||
if stride != 1 or in_planes != self.expansion*planes:
|
||||
self.downsample = [
|
||||
Conv2d(in_planes, self.expansion*planes, kernel_size=1, stride=stride, bias=False),
|
||||
BatchNorm(self.expansion*planes)
|
||||
]
|
||||
|
||||
def __call__(self, x):
|
||||
out = self.bn1(self.conv1(x)).relu()
|
||||
out = self.bn2(self.conv2(out)).relu()
|
||||
out = self.bn3(self.conv3(out))
|
||||
out = out + x.sequential(self.downsample)
|
||||
out = out.relu()
|
||||
return out
|
||||
|
||||
class ResNet:
|
||||
def __init__(self, num, num_classes=None, groups=1, width_per_group=64, stride_in_1x1=False):
|
||||
self.num = num
|
||||
self.block = {
|
||||
18: BasicBlock,
|
||||
34: BasicBlock,
|
||||
50: Bottleneck,
|
||||
101: Bottleneck,
|
||||
152: Bottleneck
|
||||
}[num]
|
||||
|
||||
self.num_blocks = {
|
||||
18: [2,2,2,2],
|
||||
34: [3,4,6,3],
|
||||
50: [3,4,6,3],
|
||||
101: [3,4,23,3],
|
||||
152: [3,8,36,3]
|
||||
}[num]
|
||||
|
||||
self.in_planes = 64
|
||||
|
||||
self.groups = groups
|
||||
self.base_width = width_per_group
|
||||
self.conv1 = Conv2d(3, 64, kernel_size=7, stride=2, bias=False, padding=3)
|
||||
self.bn1 = BatchNorm(64)
|
||||
self.layer1 = self._make_layer(self.block, 64, self.num_blocks[0], stride=1, stride_in_1x1=stride_in_1x1)
|
||||
self.layer2 = self._make_layer(self.block, 128, self.num_blocks[1], stride=2, stride_in_1x1=stride_in_1x1)
|
||||
self.layer3 = self._make_layer(self.block, 256, self.num_blocks[2], stride=2, stride_in_1x1=stride_in_1x1)
|
||||
self.layer4 = self._make_layer(self.block, 512, self.num_blocks[3], stride=2, stride_in_1x1=stride_in_1x1)
|
||||
self.fc = Linear(512 * self.block.expansion, num_classes) if num_classes is not None else None
|
||||
|
||||
def _make_layer(self, block, planes, num_blocks, stride, stride_in_1x1):
|
||||
strides = [stride] + [1] * (num_blocks-1)
|
||||
layers = []
|
||||
for stride in strides:
|
||||
if block == Bottleneck:
|
||||
layers.append(block(self.in_planes, planes, stride, stride_in_1x1, self.groups, self.base_width))
|
||||
else:
|
||||
layers.append(block(self.in_planes, planes, stride, self.groups, self.base_width))
|
||||
self.in_planes = planes * block.expansion
|
||||
return layers
|
||||
|
||||
def forward(self, x):
|
||||
is_feature_only = self.fc is None
|
||||
if is_feature_only: features = []
|
||||
out = self.bn1(self.conv1(x)).relu()
|
||||
out = out.pad([1,1,1,1]).max_pool2d((3,3), 2)
|
||||
out = out.sequential(self.layer1)
|
||||
if is_feature_only: features.append(out)
|
||||
out = out.sequential(self.layer2)
|
||||
if is_feature_only: features.append(out)
|
||||
out = out.sequential(self.layer3)
|
||||
if is_feature_only: features.append(out)
|
||||
out = out.sequential(self.layer4)
|
||||
if is_feature_only: features.append(out)
|
||||
if not is_feature_only:
|
||||
out = out.mean([2,3])
|
||||
out = self.fc(out.cast(dtypes.float32))
|
||||
return out
|
||||
return features
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return self.forward(x)
|
||||
|
||||
def load_from_pretrained(self):
|
||||
model_urls = {
|
||||
(18, 1, 64): 'https://download.pytorch.org/models/resnet18-5c106cde.pth',
|
||||
(34, 1, 64): 'https://download.pytorch.org/models/resnet34-333f7ec4.pth',
|
||||
(50, 1, 64): 'https://download.pytorch.org/models/resnet50-19c8e357.pth',
|
||||
(50, 32, 4): 'https://download.pytorch.org/models/resnext50_32x4d-7cdf4587.pth',
|
||||
(101, 1, 64): 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth',
|
||||
(152, 1, 64): 'https://download.pytorch.org/models/resnet152-b121ed2d.pth',
|
||||
}
|
||||
|
||||
self.url = model_urls[(self.num, self.groups, self.base_width)]
|
||||
for k, dat in torch_load(fetch(self.url)).items():
|
||||
try:
|
||||
obj: Tensor = get_child(self, k)
|
||||
except AttributeError as e:
|
||||
if 'fc.' in k and self.fc is None:
|
||||
continue
|
||||
|
||||
raise e
|
||||
|
||||
if 'fc.' in k and obj.shape != dat.shape:
|
||||
print("skipping fully connected layer")
|
||||
continue # Skip FC if transfer learning
|
||||
|
||||
if 'bn' not in k and 'downsample' not in k: assert obj.shape == dat.shape, (k, obj.shape, dat.shape)
|
||||
obj.assign(dat.to(obj.device).cast(obj.dtype).reshape(obj.shape))
|
||||
|
||||
ResNet18 = lambda num_classes=1000: ResNet(18, num_classes=num_classes)
|
||||
ResNet34 = lambda num_classes=1000: ResNet(34, num_classes=num_classes)
|
||||
ResNet50 = lambda num_classes=1000: ResNet(50, num_classes=num_classes)
|
||||
ResNet101 = lambda num_classes=1000: ResNet(101, num_classes=num_classes)
|
||||
ResNet152 = lambda num_classes=1000: ResNet(152, num_classes=num_classes)
|
||||
ResNeXt50_32X4D = lambda num_classes=1000: ResNet(50, num_classes=num_classes, groups=32, width_per_group=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
model = ResNet18()
|
||||
model.load_from_pretrained()
|
||||
from tinygrad import Context, GlobalCounters, TinyJit
|
||||
jmodel = TinyJit(model)
|
||||
jmodel(Tensor.rand(1, 3, 224, 224)).realize()
|
||||
GlobalCounters.reset()
|
||||
jmodel(Tensor.rand(1, 3, 224, 224)).realize()
|
||||
for i in range(10): jmodel(Tensor.rand(1, 3, 224, 224))
|
||||
263
tinygrad_repo/extra/models/retinanet.py
Normal file
263
tinygrad_repo/extra/models/retinanet.py
Normal file
@@ -0,0 +1,263 @@
|
||||
import math
|
||||
from tinygrad import Tensor, dtypes
|
||||
from tinygrad.helpers import flatten, get_child
|
||||
from examples.mlperf.helpers import generate_anchors, BoxCoder
|
||||
from examples.mlperf.losses import sigmoid_focal_loss, l1_loss
|
||||
from extra.models.resnet import ResNet
|
||||
import tinygrad.nn as nn
|
||||
import numpy as np
|
||||
|
||||
ConvFPN = ConvHead = ConvClassificationHeadLogits = nn.Conv2d
|
||||
|
||||
def nms(boxes, scores, thresh=0.5):
|
||||
x1, y1, x2, y2 = np.rollaxis(boxes, 1)
|
||||
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
|
||||
to_process, keep = scores.argsort()[::-1], []
|
||||
while to_process.size > 0:
|
||||
cur, to_process = to_process[0], to_process[1:]
|
||||
keep.append(cur)
|
||||
inter_x1 = np.maximum(x1[cur], x1[to_process])
|
||||
inter_y1 = np.maximum(y1[cur], y1[to_process])
|
||||
inter_x2 = np.minimum(x2[cur], x2[to_process])
|
||||
inter_y2 = np.minimum(y2[cur], y2[to_process])
|
||||
inter_area = np.maximum(0, inter_x2 - inter_x1 + 1) * np.maximum(0, inter_y2 - inter_y1 + 1)
|
||||
iou = inter_area / (areas[cur] + areas[to_process] - inter_area)
|
||||
to_process = to_process[np.where(iou <= thresh)[0]]
|
||||
return keep
|
||||
|
||||
def decode_bbox(offsets, anchors):
|
||||
dx, dy, dw, dh = np.rollaxis(offsets, 1)
|
||||
widths, heights = anchors[:, 2] - anchors[:, 0], anchors[:, 3] - anchors[:, 1]
|
||||
cx, cy = anchors[:, 0] + 0.5 * widths, anchors[:, 1] + 0.5 * heights
|
||||
pred_cx, pred_cy = dx * widths + cx, dy * heights + cy
|
||||
pred_w, pred_h = np.exp(dw) * widths, np.exp(dh) * heights
|
||||
pred_x1, pred_y1 = pred_cx - 0.5 * pred_w, pred_cy - 0.5 * pred_h
|
||||
pred_x2, pred_y2 = pred_cx + 0.5 * pred_w, pred_cy + 0.5 * pred_h
|
||||
return np.stack([pred_x1, pred_y1, pred_x2, pred_y2], axis=1, dtype=np.float32)
|
||||
|
||||
class RetinaNet:
|
||||
def __init__(self, backbone:ResNet, num_classes:int=264, num_anchors:int=9, scales:list[int]|None=None, aspect_ratios:list[float]|None=None):
|
||||
assert isinstance(backbone, ResNet)
|
||||
scales = tuple((i, int(i*2**(1/3)), int(i*2**(2/3))) for i in 2**np.arange(5, 10)) if scales is None else scales
|
||||
aspect_ratios = ((0.5, 1.0, 2.0),) * len(scales) if aspect_ratios is None else aspect_ratios
|
||||
self.num_anchors, self.num_classes = num_anchors, num_classes
|
||||
assert len(scales) == len(aspect_ratios) and all(self.num_anchors == len(s) * len(ar) for s, ar in zip(scales, aspect_ratios))
|
||||
|
||||
self.backbone = ResNetFPN(backbone)
|
||||
self.head = RetinaHead(self.backbone.out_channels, num_anchors=num_anchors, num_classes=num_classes)
|
||||
|
||||
def __call__(self, x:Tensor, **kwargs):
|
||||
return self.forward(x, **kwargs)
|
||||
|
||||
def forward(self, x:Tensor, **kwargs):
|
||||
return self.head(self.backbone(x), **kwargs)
|
||||
|
||||
def load_from_pretrained(self):
|
||||
model_urls = {
|
||||
(50, 1, 64): "https://download.pytorch.org/models/retinanet_resnet50_fpn_coco-eeacb38b.pth",
|
||||
(50, 32, 4): "https://zenodo.org/record/6605272/files/retinanet_model_10.zip",
|
||||
}
|
||||
self.url = model_urls[(self.backbone.body.num, self.backbone.body.groups, self.backbone.body.base_width)]
|
||||
from torch.hub import load_state_dict_from_url
|
||||
state_dict = load_state_dict_from_url(self.url, progress=True, map_location='cpu')
|
||||
state_dict = state_dict['model'] if 'model' in state_dict.keys() else state_dict
|
||||
for k, v in state_dict.items():
|
||||
obj = get_child(self, k)
|
||||
dat = v.detach().numpy()
|
||||
assert obj.shape == dat.shape, (k, obj.shape, dat.shape)
|
||||
obj.assign(dat)
|
||||
|
||||
# predictions: (BS, (H1W1+...+HmWm)A, 4 + K)
|
||||
def postprocess_detections(self, predictions, input_size=(800, 800), image_sizes=None, orig_image_sizes=None, score_thresh=0.05, topk_candidates=1000, nms_thresh=0.5):
|
||||
anchors = generate_anchors(input_size)
|
||||
grid_sizes = self.backbone.compute_grid_sizes(input_size)
|
||||
split_idx = np.cumsum([int(self.num_anchors * sz[0] * sz[1]) for sz in grid_sizes[:-1]])
|
||||
detections = []
|
||||
for i, predictions_per_image in enumerate(predictions):
|
||||
h, w = input_size if image_sizes is None else image_sizes[i]
|
||||
|
||||
predictions_per_image = np.split(predictions_per_image, split_idx)
|
||||
offsets_per_image = [br[:, :4] for br in predictions_per_image]
|
||||
scores_per_image = [cl[:, 4:] for cl in predictions_per_image]
|
||||
|
||||
image_boxes, image_scores, image_labels = [], [], []
|
||||
for offsets_per_level, scores_per_level, anchors_per_level in zip(offsets_per_image, scores_per_image, anchors):
|
||||
# remove low scoring boxes
|
||||
scores_per_level = scores_per_level.flatten()
|
||||
keep_idxs = scores_per_level > score_thresh
|
||||
scores_per_level = scores_per_level[keep_idxs]
|
||||
|
||||
# keep topk
|
||||
topk_idxs = np.where(keep_idxs)[0]
|
||||
num_topk = min(len(topk_idxs), topk_candidates)
|
||||
sort_idxs = scores_per_level.argsort()[-num_topk:][::-1]
|
||||
topk_idxs, scores_per_level = topk_idxs[sort_idxs], scores_per_level[sort_idxs]
|
||||
|
||||
# bbox coords from offsets
|
||||
anchor_idxs = topk_idxs // self.num_classes
|
||||
labels_per_level = topk_idxs % self.num_classes
|
||||
boxes_per_level = decode_bbox(offsets_per_level[anchor_idxs], anchors_per_level[anchor_idxs])
|
||||
# clip to image size
|
||||
clipped_x = boxes_per_level[:, 0::2].clip(0, w)
|
||||
clipped_y = boxes_per_level[:, 1::2].clip(0, h)
|
||||
boxes_per_level = np.stack([clipped_x, clipped_y], axis=2).reshape(-1, 4)
|
||||
|
||||
image_boxes.append(boxes_per_level)
|
||||
image_scores.append(scores_per_level)
|
||||
image_labels.append(labels_per_level)
|
||||
|
||||
image_boxes = np.concatenate(image_boxes)
|
||||
image_scores = np.concatenate(image_scores)
|
||||
image_labels = np.concatenate(image_labels)
|
||||
|
||||
# nms for each class
|
||||
keep_mask = np.zeros_like(image_scores, dtype=bool)
|
||||
for class_id in np.unique(image_labels):
|
||||
curr_indices = np.where(image_labels == class_id)[0]
|
||||
curr_keep_indices = nms(image_boxes[curr_indices], image_scores[curr_indices], nms_thresh)
|
||||
keep_mask[curr_indices[curr_keep_indices]] = True
|
||||
keep = np.where(keep_mask)[0]
|
||||
keep = keep[image_scores[keep].argsort()[::-1]]
|
||||
|
||||
# resize bboxes back to original size
|
||||
image_boxes = image_boxes[keep]
|
||||
if orig_image_sizes is not None:
|
||||
resized_x = image_boxes[:, 0::2] * orig_image_sizes[i][1] / w
|
||||
resized_y = image_boxes[:, 1::2] * orig_image_sizes[i][0] / h
|
||||
image_boxes = np.stack([resized_x, resized_y], axis=2).reshape(-1, 4)
|
||||
# xywh format
|
||||
image_boxes = np.concatenate([image_boxes[:, :2], image_boxes[:, 2:] - image_boxes[:, :2]], axis=1)
|
||||
|
||||
detections.append({"boxes":image_boxes, "scores":image_scores[keep], "labels":image_labels[keep]})
|
||||
return detections
|
||||
|
||||
class ClassificationHead:
|
||||
def __init__(self, in_channels:int, num_anchors:int, num_classes:int):
|
||||
self.num_classes = num_classes
|
||||
self.conv = flatten([(ConvHead(in_channels, in_channels, kernel_size=3, padding=1), lambda x: x.relu()) for _ in range(4)])
|
||||
self.cls_logits = ConvClassificationHeadLogits(in_channels, num_anchors * num_classes, kernel_size=3, padding=1)
|
||||
|
||||
def __call__(self, x:Tensor, labels:Tensor|None=None, matches:Tensor|None=None):
|
||||
out = [self.cls_logits(feat.sequential(self.conv)).permute(0, 2, 3, 1).reshape(feat.shape[0], -1, self.num_classes) for feat in x]
|
||||
out = out[0].cat(*out[1:], dim=1)
|
||||
|
||||
if Tensor.training:
|
||||
assert labels is not None and matches is not None, "labels and matches should be passed in when training"
|
||||
return self._compute_loss(out.cast(dtypes.float32), labels, matches)
|
||||
|
||||
return out.sigmoid()
|
||||
|
||||
def _compute_loss(self, x:Tensor, labels:Tensor, matches:Tensor) -> Tensor:
|
||||
labels = ((labels + 1) * (fg_idxs := matches >= 0) - 1).one_hot(num_classes=x.shape[-1])
|
||||
valid_idxs = (matches != -2).reshape(matches.shape[0], -1, 1)
|
||||
loss = valid_idxs.where(sigmoid_focal_loss(x, labels), 0).sum((-1, -2))
|
||||
loss = (loss / fg_idxs.sum(-1)).sum() / matches.shape[0]
|
||||
return loss
|
||||
|
||||
class RegressionHead:
|
||||
def __init__(self, in_channels:int, num_anchors:int, box_coder:BoxCoder|None=None):
|
||||
self.conv = flatten([(ConvHead(in_channels, in_channels, kernel_size=3, padding=1), lambda x: x.relu()) for _ in range(4)])
|
||||
self.bbox_reg = ConvHead(in_channels, num_anchors * 4, kernel_size=3, padding=1)
|
||||
|
||||
if box_coder is None:
|
||||
box_coder = BoxCoder((1.0, 1.0, 1.0, 1.0))
|
||||
self.box_coder = box_coder
|
||||
|
||||
def __call__(self, x:Tensor, bboxes:Tensor|None=None, matches:Tensor|None=None, anchors:Tensor|None=None):
|
||||
out = [self.bbox_reg(feat.sequential(self.conv)).permute(0, 2, 3, 1).reshape(feat.shape[0], -1, 4) for feat in x]
|
||||
out = out[0].cat(*out[1:], dim=1)
|
||||
|
||||
if Tensor.training:
|
||||
assert bboxes is not None and matches is not None and anchors is not None, "bboxes, matches, and anchors should be passed in when training"
|
||||
return self._compute_loss(out, bboxes, matches, anchors)
|
||||
|
||||
return out
|
||||
|
||||
def _compute_loss(self, x:Tensor, bboxes:Tensor, matches:Tensor, anchors:Tensor) -> Tensor:
|
||||
mask = (fg_idxs := matches >= 0).reshape(matches.shape[0], -1, 1)
|
||||
x = x * mask
|
||||
tgt = self.box_coder.encode(bboxes, anchors) * mask
|
||||
loss = l1_loss(x, tgt).sum((-1, -2))
|
||||
loss = (loss / fg_idxs.sum(-1)).sum() / matches.shape[0]
|
||||
return loss
|
||||
|
||||
class RetinaHead:
|
||||
def __init__(self, in_channels:int, num_anchors:int, num_classes:int):
|
||||
self.classification_head = ClassificationHead(in_channels, num_anchors, num_classes)
|
||||
self.regression_head = RegressionHead(in_channels, num_anchors)
|
||||
|
||||
def __call__(self, x:Tensor, **kwargs) -> Tensor|dict[str, Tensor]:
|
||||
if Tensor.training:
|
||||
return {
|
||||
"classification_loss": self.classification_head(x, labels=kwargs["labels"], matches=kwargs["matches"]),
|
||||
"regression_loss": self.regression_head(x, bboxes=kwargs["bboxes"], matches=kwargs["matches"], anchors=kwargs["anchors"])
|
||||
}
|
||||
|
||||
pred_bbox, pred_class = self.regression_head(x), self.classification_head(x)
|
||||
out = pred_bbox.cat(pred_class, dim=-1)
|
||||
return out
|
||||
|
||||
class ResNetFPN:
|
||||
def __init__(self, resnet:ResNet, out_channels:int=256, returned_layers:list[int]=[2, 3, 4]):
|
||||
self.out_channels = out_channels
|
||||
self.body = resnet
|
||||
in_channels_list = [(self.body.in_planes // 8) * 2 ** (i - 1) for i in returned_layers]
|
||||
self.fpn = FPN(in_channels_list, out_channels)
|
||||
|
||||
# this is needed to decouple inference from postprocessing (anchors generation)
|
||||
def compute_grid_sizes(self, input_size):
|
||||
return np.ceil(np.array(input_size)[None, :] / 2 ** np.arange(3, 8)[:, None])
|
||||
|
||||
def __call__(self, x:Tensor):
|
||||
out = self.body.bn1(self.body.conv1(x)).relu()
|
||||
out = out.pad([1,1,1,1]).max_pool2d((3,3), 2)
|
||||
out = out.sequential(self.body.layer1)
|
||||
p3 = out.sequential(self.body.layer2)
|
||||
p4 = p3.sequential(self.body.layer3)
|
||||
p5 = p4.sequential(self.body.layer4)
|
||||
return self.fpn([p3, p4, p5])
|
||||
|
||||
class ExtraFPNBlock:
|
||||
def __init__(self, in_channels:int, out_channels:int):
|
||||
self.p6 = ConvFPN(in_channels, out_channels, kernel_size=3, stride=2, padding=1)
|
||||
self.p7 = ConvFPN(out_channels, out_channels, kernel_size=3, stride=2, padding=1)
|
||||
self.use_P5 = in_channels == out_channels
|
||||
|
||||
def __call__(self, p:Tensor, c:Tensor):
|
||||
p5, c5 = p[-1], c[-1]
|
||||
x = p5 if self.use_P5 else c5
|
||||
p6 = self.p6(x)
|
||||
p7 = self.p7(p6.relu())
|
||||
p.extend([p6, p7])
|
||||
return p
|
||||
|
||||
class FPN:
|
||||
def __init__(self, in_channels_list:list[int], out_channels:int, extra_blocks:ExtraFPNBlock|None=None):
|
||||
self.inner_blocks, self.layer_blocks = [], []
|
||||
for in_channels in in_channels_list:
|
||||
self.inner_blocks.append(ConvFPN(in_channels, out_channels, kernel_size=1))
|
||||
self.layer_blocks.append(ConvFPN(out_channels, out_channels, kernel_size=3, padding=1))
|
||||
self.extra_blocks = ExtraFPNBlock(256, 256) if extra_blocks is None else extra_blocks
|
||||
|
||||
def __call__(self, x:Tensor):
|
||||
last_inner = self.inner_blocks[-1](x[-1])
|
||||
results = [self.layer_blocks[-1](last_inner)]
|
||||
for idx in range(len(x) - 2, -1, -1):
|
||||
inner_lateral = self.inner_blocks[idx](x[idx])
|
||||
|
||||
# upsample to inner_lateral's shape
|
||||
(ih, iw), (oh, ow), prefix = last_inner.shape[-2:], inner_lateral.shape[-2:], last_inner.shape[:-2]
|
||||
eh, ew = math.ceil(oh / ih), math.ceil(ow / iw)
|
||||
inner_top_down = last_inner.reshape(*prefix, ih, 1, iw, 1).expand(*prefix, ih, eh, iw, ew).reshape(*prefix, ih*eh, iw*ew)[:, :, :oh, :ow]
|
||||
|
||||
last_inner = inner_lateral + inner_top_down
|
||||
results.insert(0, self.layer_blocks[idx](last_inner))
|
||||
if self.extra_blocks is not None:
|
||||
results = self.extra_blocks(results, x)
|
||||
return results
|
||||
|
||||
if __name__ == "__main__":
|
||||
from extra.models.resnet import ResNeXt50_32X4D
|
||||
backbone = ResNeXt50_32X4D()
|
||||
retina = RetinaNet(backbone)
|
||||
retina.load_from_pretrained()
|
||||
202
tinygrad_repo/extra/models/rnnt.py
Normal file
202
tinygrad_repo/extra/models/rnnt.py
Normal file
@@ -0,0 +1,202 @@
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.engine.jit import TinyJit
|
||||
from tinygrad.nn import Linear, Embedding
|
||||
from tinygrad.helpers import fetch
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class RNNT:
|
||||
def __init__(self, input_features=240, vocab_size=29, enc_hidden_size=1024, pred_hidden_size=320, joint_hidden_size=512, pre_enc_layers=2, post_enc_layers=3, pred_layers=2, stack_time_factor=2, dropout=0.32):
|
||||
self.encoder = Encoder(input_features, enc_hidden_size, pre_enc_layers, post_enc_layers, stack_time_factor, dropout)
|
||||
self.prediction = Prediction(vocab_size, pred_hidden_size, pred_layers, dropout)
|
||||
self.joint = Joint(vocab_size, pred_hidden_size, enc_hidden_size, joint_hidden_size, dropout)
|
||||
|
||||
@TinyJit
|
||||
def __call__(self, x, y, hc=None):
|
||||
f, _ = self.encoder(x, None)
|
||||
g, _ = self.prediction(y, hc, Tensor.ones(1))
|
||||
out = self.joint(f, g)
|
||||
return out.realize()
|
||||
|
||||
def decode(self, x, x_lens):
|
||||
logits, logit_lens = self.encoder(x, x_lens)
|
||||
outputs = []
|
||||
for b in range(logits.shape[0]):
|
||||
inseq = logits[b, :, :].unsqueeze(1)
|
||||
logit_len = logit_lens[b]
|
||||
seq = self._greedy_decode(inseq, int(np.ceil(logit_len.numpy()).item()))
|
||||
outputs.append(seq)
|
||||
return outputs
|
||||
|
||||
def _greedy_decode(self, logits, logit_len):
|
||||
hc = Tensor.zeros(self.prediction.rnn.layers, 2, self.prediction.hidden_size)
|
||||
labels = []
|
||||
label = Tensor.zeros(1, 1)
|
||||
mask = Tensor.zeros(1)
|
||||
for time_idx in range(logit_len):
|
||||
logit = logits[time_idx, :, :].unsqueeze(0)
|
||||
not_blank = True
|
||||
added = 0
|
||||
while not_blank and added < 30:
|
||||
if len(labels) > 0:
|
||||
mask = (mask + 1).clip(0, 1)
|
||||
label = Tensor([[labels[-1] if labels[-1] <= 28 else labels[-1] - 1]]) + 1 - 1
|
||||
jhc = self._pred_joint(Tensor(logit.numpy()), label, hc, mask)
|
||||
k = jhc[0, 0, :29].argmax(axis=0).numpy()
|
||||
not_blank = k != 28
|
||||
if not_blank:
|
||||
labels.append(k)
|
||||
hc = jhc[:, :, 29:] + 1 - 1
|
||||
added += 1
|
||||
return labels
|
||||
|
||||
@TinyJit
|
||||
def _pred_joint(self, logit, label, hc, mask):
|
||||
g, hc = self.prediction(label, hc, mask)
|
||||
j = self.joint(logit, g)[0]
|
||||
j = j.pad(((0, 1), (0, 1), (0, 0)))
|
||||
out = j.cat(hc, dim=2)
|
||||
return out.realize()
|
||||
|
||||
def load_from_pretrained(self):
|
||||
fn = Path(__file__).parents[1] / "weights/rnnt.pt"
|
||||
fetch("https://zenodo.org/record/3662521/files/DistributedDataParallel_1576581068.9962234-epoch-100.pt?download=1", fn)
|
||||
|
||||
import torch
|
||||
with open(fn, "rb") as f:
|
||||
state_dict = torch.load(f, map_location="cpu")["state_dict"]
|
||||
|
||||
# encoder
|
||||
for i in range(2):
|
||||
self.encoder.pre_rnn.cells[i].weights_ih.assign(state_dict[f"encoder.pre_rnn.lstm.weight_ih_l{i}"].numpy())
|
||||
self.encoder.pre_rnn.cells[i].weights_hh.assign(state_dict[f"encoder.pre_rnn.lstm.weight_hh_l{i}"].numpy())
|
||||
self.encoder.pre_rnn.cells[i].bias_ih.assign(state_dict[f"encoder.pre_rnn.lstm.bias_ih_l{i}"].numpy())
|
||||
self.encoder.pre_rnn.cells[i].bias_hh.assign(state_dict[f"encoder.pre_rnn.lstm.bias_hh_l{i}"].numpy())
|
||||
for i in range(3):
|
||||
self.encoder.post_rnn.cells[i].weights_ih.assign(state_dict[f"encoder.post_rnn.lstm.weight_ih_l{i}"].numpy())
|
||||
self.encoder.post_rnn.cells[i].weights_hh.assign(state_dict[f"encoder.post_rnn.lstm.weight_hh_l{i}"].numpy())
|
||||
self.encoder.post_rnn.cells[i].bias_ih.assign(state_dict[f"encoder.post_rnn.lstm.bias_ih_l{i}"].numpy())
|
||||
self.encoder.post_rnn.cells[i].bias_hh.assign(state_dict[f"encoder.post_rnn.lstm.bias_hh_l{i}"].numpy())
|
||||
|
||||
# prediction
|
||||
self.prediction.emb.weight.assign(state_dict["prediction.embed.weight"].numpy())
|
||||
for i in range(2):
|
||||
self.prediction.rnn.cells[i].weights_ih.assign(state_dict[f"prediction.dec_rnn.lstm.weight_ih_l{i}"].numpy())
|
||||
self.prediction.rnn.cells[i].weights_hh.assign(state_dict[f"prediction.dec_rnn.lstm.weight_hh_l{i}"].numpy())
|
||||
self.prediction.rnn.cells[i].bias_ih.assign(state_dict[f"prediction.dec_rnn.lstm.bias_ih_l{i}"].numpy())
|
||||
self.prediction.rnn.cells[i].bias_hh.assign(state_dict[f"prediction.dec_rnn.lstm.bias_hh_l{i}"].numpy())
|
||||
|
||||
# joint
|
||||
self.joint.l1.weight.assign(state_dict["joint_net.0.weight"].numpy())
|
||||
self.joint.l1.bias.assign(state_dict["joint_net.0.bias"].numpy())
|
||||
self.joint.l2.weight.assign(state_dict["joint_net.3.weight"].numpy())
|
||||
self.joint.l2.bias.assign(state_dict["joint_net.3.bias"].numpy())
|
||||
|
||||
|
||||
class LSTMCell:
|
||||
def __init__(self, input_size, hidden_size, dropout):
|
||||
self.dropout = dropout
|
||||
|
||||
self.weights_ih = Tensor.uniform(hidden_size * 4, input_size)
|
||||
self.bias_ih = Tensor.uniform(hidden_size * 4)
|
||||
self.weights_hh = Tensor.uniform(hidden_size * 4, hidden_size)
|
||||
self.bias_hh = Tensor.uniform(hidden_size * 4)
|
||||
|
||||
def __call__(self, x, hc):
|
||||
gates = x.linear(self.weights_ih.T, self.bias_ih) + hc[:x.shape[0]].linear(self.weights_hh.T, self.bias_hh)
|
||||
|
||||
i, f, g, o = gates.chunk(4, 1)
|
||||
i, f, g, o = i.sigmoid(), f.sigmoid(), g.tanh(), o.sigmoid()
|
||||
|
||||
c = (f * hc[x.shape[0]:]) + (i * g)
|
||||
h = (o * c.tanh()).dropout(self.dropout)
|
||||
|
||||
return Tensor.cat(h, c).realize()
|
||||
|
||||
|
||||
class LSTM:
|
||||
def __init__(self, input_size, hidden_size, layers, dropout):
|
||||
self.input_size = input_size
|
||||
self.hidden_size = hidden_size
|
||||
self.layers = layers
|
||||
|
||||
self.cells = [LSTMCell(input_size, hidden_size, dropout) if i == 0 else LSTMCell(hidden_size, hidden_size, dropout if i != layers - 1 else 0) for i in range(layers)]
|
||||
|
||||
def __call__(self, x, hc):
|
||||
@TinyJit
|
||||
def _do_step(x_, hc_):
|
||||
return self.do_step(x_, hc_)
|
||||
|
||||
if hc is None:
|
||||
hc = Tensor.zeros(self.layers, 2 * x.shape[1], self.hidden_size).contiguous().realize()
|
||||
|
||||
output = None
|
||||
for t in range(x.shape[0]):
|
||||
hc = _do_step(x[t] + 1 - 1, hc) # TODO: why do we need to do this?
|
||||
if output is None:
|
||||
output = hc[-1:, :x.shape[1]]
|
||||
else:
|
||||
output = output.cat(hc[-1:, :x.shape[1]], dim=0).realize()
|
||||
|
||||
return output, hc
|
||||
|
||||
def do_step(self, x, hc):
|
||||
new_hc = [x]
|
||||
for i, cell in enumerate(self.cells):
|
||||
new_hc.append(cell(new_hc[i][:x.shape[0]], hc[i]))
|
||||
return Tensor.stack(*new_hc[1:]).realize()
|
||||
|
||||
|
||||
class StackTime:
|
||||
def __init__(self, factor):
|
||||
self.factor = factor
|
||||
|
||||
def __call__(self, x, x_lens):
|
||||
x = x.pad(((0, (-x.shape[0]) % self.factor), (0, 0), (0, 0)))
|
||||
x = x.reshape(x.shape[0] // self.factor, x.shape[1], x.shape[2] * self.factor)
|
||||
return x, x_lens / self.factor if x_lens is not None else None
|
||||
|
||||
|
||||
class Encoder:
|
||||
def __init__(self, input_size, hidden_size, pre_layers, post_layers, stack_time_factor, dropout):
|
||||
self.pre_rnn = LSTM(input_size, hidden_size, pre_layers, dropout)
|
||||
self.stack_time = StackTime(stack_time_factor)
|
||||
self.post_rnn = LSTM(stack_time_factor * hidden_size, hidden_size, post_layers, dropout)
|
||||
|
||||
def __call__(self, x, x_lens):
|
||||
x, _ = self.pre_rnn(x, None)
|
||||
x, x_lens = self.stack_time(x, x_lens)
|
||||
x, _ = self.post_rnn(x, None)
|
||||
return x.transpose(0, 1), x_lens
|
||||
|
||||
|
||||
class Prediction:
|
||||
def __init__(self, vocab_size, hidden_size, layers, dropout):
|
||||
self.hidden_size = hidden_size
|
||||
|
||||
self.emb = Embedding(vocab_size - 1, hidden_size)
|
||||
self.rnn = LSTM(hidden_size, hidden_size, layers, dropout)
|
||||
|
||||
def __call__(self, x, hc, m):
|
||||
emb = self.emb(x) * m
|
||||
x_, hc = self.rnn(emb.transpose(0, 1), hc)
|
||||
return x_.transpose(0, 1), hc
|
||||
|
||||
|
||||
class Joint:
|
||||
def __init__(self, vocab_size, pred_hidden_size, enc_hidden_size, joint_hidden_size, dropout):
|
||||
self.dropout = dropout
|
||||
|
||||
self.l1 = Linear(pred_hidden_size + enc_hidden_size, joint_hidden_size)
|
||||
self.l2 = Linear(joint_hidden_size, vocab_size)
|
||||
|
||||
def __call__(self, f, g):
|
||||
(_, T, H), (B, U, H2) = f.shape, g.shape
|
||||
f = f.unsqueeze(2).expand(B, T, U, H)
|
||||
g = g.unsqueeze(1).expand(B, T, U, H2)
|
||||
|
||||
inp = f.cat(g, dim=3)
|
||||
t = self.l1(inp).relu()
|
||||
t = t.dropout(self.dropout)
|
||||
return self.l2(t)
|
||||
288
tinygrad_repo/extra/models/t5.py
Normal file
288
tinygrad_repo/extra/models/t5.py
Normal file
@@ -0,0 +1,288 @@
|
||||
# pip3 install sentencepiece
|
||||
|
||||
# adapted from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/modeling_t5.py
|
||||
|
||||
# coding=utf-8
|
||||
# Copyright 2018 Mesh TensorFlow authors, T5 Authors and HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Tinygrad T5 model."""
|
||||
from tinygrad import nn, Tensor, dtypes
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Union, Optional, Tuple
|
||||
from pathlib import Path
|
||||
from sentencepiece import SentencePieceProcessor
|
||||
|
||||
# default config is t5-xxl
|
||||
@dataclass
|
||||
class T5Config:
|
||||
d_ff:int = 10240
|
||||
d_kv:int = 64
|
||||
d_model:int = 4096
|
||||
layer_norm_epsilon:float = 1e-6
|
||||
num_decoder_layers:int = 24
|
||||
num_heads:int = 64
|
||||
num_layers:int = 24
|
||||
relative_attention_num_buckets:int = 32
|
||||
relative_attention_max_distance:int = 128
|
||||
vocab_size:int = 32128
|
||||
|
||||
class T5Tokenizer:
|
||||
def __init__(self, spiece_path):
|
||||
self.spp = SentencePieceProcessor(str(spiece_path))
|
||||
|
||||
def encode(self, text:str, max_length:int) -> List[int]:
|
||||
encoded = self.spp.Encode(text)
|
||||
if len(encoded) > max_length - 1: encoded = encoded[:max_length - 1]
|
||||
return encoded + [1] + [0]*(max_length - len(encoded) - 1)
|
||||
|
||||
class T5LayerNorm:
|
||||
def __init__(self, hidden_size:int, eps:float=1e-6):
|
||||
"""
|
||||
Construct a layernorm module in the T5 style. No bias and no subtraction of mean.
|
||||
"""
|
||||
self.weight = Tensor.ones(hidden_size)
|
||||
self.variance_epsilon = eps
|
||||
|
||||
def __call__(self, hidden_states:Tensor) -> Tensor:
|
||||
# T5 uses a layer_norm which only scales and doesn't shift, which is also known as Root Mean
|
||||
# Square Layer Normalization https://arxiv.org/abs/1910.07467 thus varience is calculated
|
||||
# w/o mean and there is no bias. Additionally we want to make sure that the accumulation for
|
||||
# half-precision inputs is done in fp32
|
||||
|
||||
variance = hidden_states.cast(dtypes.float32).pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * Tensor.rsqrt(variance + self.variance_epsilon)
|
||||
|
||||
# convert into half-precision if necessary
|
||||
if self.weight.dtype in [dtypes.float16, dtypes.bfloat16]:
|
||||
hidden_states = hidden_states.cast(self.weight.dtype)
|
||||
|
||||
return self.weight * hidden_states
|
||||
|
||||
|
||||
class T5DenseGatedActDense:
|
||||
def __init__(self, config:T5Config):
|
||||
self.wi_0 = nn.Linear(config.d_model, config.d_ff, bias=False)
|
||||
self.wi_1 = nn.Linear(config.d_model, config.d_ff, bias=False)
|
||||
self.wo = nn.Linear(config.d_ff, config.d_model, bias=False)
|
||||
|
||||
def __call__(self, hidden_states:Tensor) -> Tensor:
|
||||
hidden_gelu = self.wi_0(hidden_states).gelu()
|
||||
hidden_linear = self.wi_1(hidden_states)
|
||||
hidden_states = hidden_gelu * hidden_linear
|
||||
hidden_states = self.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5LayerFF:
|
||||
def __init__(self, config:T5Config):
|
||||
self.DenseReluDense = T5DenseGatedActDense(config)
|
||||
self.layer_norm = T5LayerNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
|
||||
def __call__(self, hidden_states:Tensor) -> Tensor:
|
||||
forwarded_states = self.layer_norm(hidden_states)
|
||||
forwarded_states = self.DenseReluDense(forwarded_states)
|
||||
hidden_states = hidden_states + forwarded_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Attention:
|
||||
def __init__(self, config:T5Config, has_relative_attention_bias:bool=False):
|
||||
self.has_relative_attention_bias = has_relative_attention_bias
|
||||
self.relative_attention_num_buckets = config.relative_attention_num_buckets
|
||||
self.relative_attention_max_distance = config.relative_attention_max_distance
|
||||
self.d_model = config.d_model
|
||||
self.key_value_proj_dim = config.d_kv
|
||||
self.n_heads = config.num_heads
|
||||
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
||||
|
||||
# Mesh TensorFlow initialization to avoid scaling before softmax
|
||||
self.q = nn.Linear(self.d_model, self.inner_dim, bias=False)
|
||||
self.k = nn.Linear(self.d_model, self.inner_dim, bias=False)
|
||||
self.v = nn.Linear(self.d_model, self.inner_dim, bias=False)
|
||||
self.o = nn.Linear(self.inner_dim, self.d_model, bias=False)
|
||||
|
||||
if self.has_relative_attention_bias:
|
||||
self.relative_attention_bias = nn.Embedding(self.relative_attention_num_buckets, self.n_heads)
|
||||
|
||||
@staticmethod
|
||||
def _relative_position_bucket(relative_position:Tensor, num_buckets:int=32, max_distance:int=128) -> Tensor:
|
||||
"""
|
||||
Adapted from Mesh Tensorflow:
|
||||
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
|
||||
|
||||
Translate relative position to a bucket number for relative attention. The relative position is defined as
|
||||
memory_position - query_position, i.e. the distance in tokens from the attending position to the attended-to
|
||||
position. If bidirectional=False, then positive relative positions are invalid. We use smaller buckets for
|
||||
small absolute relative_position and larger buckets for larger absolute relative_positions. All relative
|
||||
positions >=max_distance map to the same bucket. All relative positions <=-max_distance map to the same bucket.
|
||||
This should allow for more graceful generalization to longer sequences than the model has been trained on
|
||||
|
||||
Args:
|
||||
relative_position: an int32 Tensor
|
||||
bidirectional: a boolean - whether the attention is bidirectional
|
||||
num_buckets: an integer
|
||||
max_distance: an integer
|
||||
|
||||
Returns:
|
||||
a Tensor with the same shape as relative_position, containing int32 values in the range [0, num_buckets)
|
||||
"""
|
||||
relative_buckets = Tensor.zeros_like(relative_position)
|
||||
num_buckets //= 2
|
||||
relative_buckets += (relative_position > 0).cast(dtypes.long) * num_buckets
|
||||
relative_position = Tensor.abs(relative_position)
|
||||
|
||||
# half of the buckets are for exact increments in positions
|
||||
max_exact = num_buckets // 2
|
||||
is_small = relative_position < max_exact
|
||||
|
||||
# The other half of the buckets are for logarithmically bigger bins in positions up to max_distance
|
||||
relative_position_if_large = max_exact + (
|
||||
Tensor.log(relative_position.float() / max_exact)
|
||||
/ math.log(max_distance / max_exact)
|
||||
* (num_buckets - max_exact)
|
||||
).cast(dtypes.long)
|
||||
|
||||
relative_position_if_large = Tensor.min(
|
||||
Tensor.stack(
|
||||
relative_position_if_large, Tensor.full_like(relative_position_if_large, num_buckets - 1)
|
||||
),
|
||||
axis=0,
|
||||
)
|
||||
relative_buckets += Tensor.where(is_small, relative_position, relative_position_if_large)
|
||||
return relative_buckets
|
||||
|
||||
def compute_bias(self, query_length, key_length, device=None) -> Tensor:
|
||||
"""Compute binned relative position bias"""
|
||||
if device is None:
|
||||
device = self.relative_attention_bias.weight.device
|
||||
context_position = Tensor.arange(query_length, dtype=dtypes.long, device=device)[:, None]
|
||||
memory_position = Tensor.arange(key_length, dtype=dtypes.long, device=device)[None, :]
|
||||
relative_position = memory_position - context_position # shape (query_length, key_length)
|
||||
relative_position_bucket = self._relative_position_bucket(
|
||||
relative_position, # shape (query_length, key_length)
|
||||
num_buckets=self.relative_attention_num_buckets,
|
||||
max_distance=self.relative_attention_max_distance,
|
||||
)
|
||||
values = self.relative_attention_bias(relative_position_bucket) # shape (query_length, key_length, num_heads)
|
||||
values = values.permute([2, 0, 1]).unsqueeze(0) # shape (1, num_heads, query_length, key_length)
|
||||
return values
|
||||
|
||||
def __call__(self, hidden_states:Tensor, position_bias:Optional[Tensor]=None) -> Tuple[Tensor,Tensor]:
|
||||
"""
|
||||
Self-attention (if key_value_states is None) or attention over source sentence (provided by key_value_states).
|
||||
"""
|
||||
# Input is (batch_size, seq_length, dim)
|
||||
batch_size, key_length = hidden_states.shape[:2]
|
||||
|
||||
def shape(states):
|
||||
"""projection"""
|
||||
return states.view(batch_size, -1, self.n_heads, self.key_value_proj_dim).transpose(1, 2)
|
||||
|
||||
def unshape(states):
|
||||
"""reshape"""
|
||||
return states.transpose(1, 2).contiguous().view(batch_size, -1, self.inner_dim)
|
||||
|
||||
def project(hidden_states, proj_layer):
|
||||
"""projects hidden states correctly to key/query states"""
|
||||
# self-attn
|
||||
# (batch_size, n_heads, seq_length, dim_per_head)
|
||||
return shape(proj_layer(hidden_states))
|
||||
|
||||
# get query states
|
||||
query_states = shape(self.q(hidden_states)) # (batch_size, n_heads, seq_length, dim_per_head)
|
||||
|
||||
# get key/value states
|
||||
key_states = project(hidden_states, self.k)
|
||||
value_states = project(hidden_states, self.v)
|
||||
|
||||
# compute scores
|
||||
scores = Tensor.matmul(query_states, key_states.transpose(3, 2))
|
||||
|
||||
if position_bias is None:
|
||||
position_bias = self.compute_bias(key_length, key_length, device=scores.device)
|
||||
|
||||
scores += position_bias
|
||||
attn_weights = Tensor.softmax(scores.float(), axis=-1).cast(scores.dtype) # (batch_size, n_heads, seq_length, key_length)
|
||||
|
||||
attn_output = unshape(Tensor.matmul(attn_weights, value_states)) # (batch_size, seq_length, dim)
|
||||
attn_output = self.o(attn_output)
|
||||
|
||||
return attn_output, position_bias
|
||||
|
||||
|
||||
class T5LayerSelfAttention:
|
||||
def __init__(self, config:T5Config, has_relative_attention_bias:bool=False):
|
||||
self.SelfAttention = T5Attention(config, has_relative_attention_bias=has_relative_attention_bias)
|
||||
self.layer_norm = T5LayerNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
|
||||
def __call__(self, hidden_states:Tensor, position_bias:Optional[Tensor]=None) -> Tuple[Tensor, Tensor]:
|
||||
normed_hidden_states = self.layer_norm(hidden_states)
|
||||
attention_output, position_bias = self.SelfAttention(normed_hidden_states, position_bias=position_bias)
|
||||
return hidden_states + attention_output, position_bias
|
||||
|
||||
|
||||
class T5Block:
|
||||
def __init__(self, config:T5Config, has_relative_attention_bias:bool=False):
|
||||
self.layer = (T5LayerSelfAttention(config, has_relative_attention_bias=has_relative_attention_bias),
|
||||
T5LayerFF(config))
|
||||
|
||||
def __call__(self, hidden_states:Tensor, position_bias:Optional[Tensor]=None) -> Tuple[Tensor, Tensor]:
|
||||
self_attention_outputs, position_bias = self.layer[0](hidden_states, position_bias=position_bias)
|
||||
hidden_states = self_attention_outputs
|
||||
|
||||
# Apply Feed Forward layer
|
||||
hidden_states = self.layer[-1](hidden_states)
|
||||
|
||||
return hidden_states, position_bias
|
||||
|
||||
|
||||
class T5Stack:
|
||||
def __init__(self, config:T5Config, embed_tokens:nn.Embedding):
|
||||
self.config = config
|
||||
self.embed_tokens = embed_tokens
|
||||
self.block = [T5Block(config, has_relative_attention_bias=bool(i == 0)) for i in range(config.num_layers)]
|
||||
self.final_layer_norm = T5LayerNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
|
||||
def __call__(self, input_ids:Tensor) -> Tensor:
|
||||
input_ids = input_ids.view(-1, input_ids.shape[-1])
|
||||
|
||||
hidden_states, position_bias = self.embed_tokens(input_ids), None
|
||||
|
||||
for layer_module in self.block:
|
||||
hidden_states, position_bias = layer_module(hidden_states, position_bias=position_bias)
|
||||
|
||||
return self.final_layer_norm(hidden_states)
|
||||
|
||||
|
||||
class T5EncoderModel:
|
||||
def __init__(self, config:T5Config):
|
||||
self.shared = nn.Embedding(config.vocab_size, config.d_model)
|
||||
self.encoder = T5Stack(config, self.shared)
|
||||
|
||||
def __call__(self, input_ids:Tensor) -> Tensor:
|
||||
return self.encoder(input_ids)
|
||||
|
||||
class T5Embedder:
|
||||
def __init__(self, max_length:int, spiece_path:Union[str, Path]):
|
||||
self.tokenizer = T5Tokenizer(spiece_path)
|
||||
self.max_length = max_length
|
||||
config = T5Config()
|
||||
self.encoder = T5EncoderModel(config)
|
||||
|
||||
def __call__(self, texts:Union[str, List[str]]) -> Tensor:
|
||||
if isinstance(texts, str): texts = [texts]
|
||||
toks = Tensor.cat(*[Tensor(self.tokenizer.encode(text, self.max_length)) for text in texts], dim=0)
|
||||
return self.encoder(toks)
|
||||
61
tinygrad_repo/extra/models/transformer.py
Normal file
61
tinygrad_repo/extra/models/transformer.py
Normal file
@@ -0,0 +1,61 @@
|
||||
from tinygrad import Tensor
|
||||
|
||||
class TransformerBlock:
|
||||
def __init__(self, embed_dim, num_heads, ff_dim, prenorm=False, act=lambda x: x.relu(), dropout=0.1):
|
||||
assert embed_dim % num_heads == 0, "embed_dim must be divisible by num_heads"
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.head_size = embed_dim // num_heads
|
||||
self.prenorm, self.act = prenorm, act
|
||||
self.dropout = dropout
|
||||
|
||||
self.query = (Tensor.scaled_uniform(embed_dim, embed_dim), Tensor.zeros(embed_dim))
|
||||
self.key = (Tensor.scaled_uniform(embed_dim, embed_dim), Tensor.zeros(embed_dim))
|
||||
self.value = (Tensor.scaled_uniform(embed_dim, embed_dim), Tensor.zeros(embed_dim))
|
||||
|
||||
self.out = (Tensor.scaled_uniform(embed_dim, embed_dim), Tensor.zeros(embed_dim))
|
||||
|
||||
self.ff1 = (Tensor.scaled_uniform(embed_dim, ff_dim), Tensor.zeros(ff_dim))
|
||||
self.ff2 = (Tensor.scaled_uniform(ff_dim, embed_dim), Tensor.zeros(embed_dim))
|
||||
|
||||
self.ln1 = (Tensor.ones(embed_dim), Tensor.zeros(embed_dim))
|
||||
self.ln2 = (Tensor.ones(embed_dim), Tensor.zeros(embed_dim))
|
||||
|
||||
def attn(self, x):
|
||||
# x: (bs, time, embed_dim) -> (bs, time, embed_dim)
|
||||
query, key, value = [x.linear(*y).reshape(shape=(x.shape[0], -1, self.num_heads, self.head_size)).transpose(1,2) for y in [self.query, self.key, self.value]]
|
||||
attention = Tensor.scaled_dot_product_attention(query, key, value).transpose(1,2)
|
||||
return attention.reshape(shape=(x.shape[0], -1, self.num_heads * self.head_size)).linear(*self.out)
|
||||
|
||||
def __call__(self, x):
|
||||
if self.prenorm:
|
||||
x = x + self.attn(x.layernorm().linear(*self.ln1)).dropout(self.dropout)
|
||||
x = x + self.act(x.layernorm().linear(*self.ln2).linear(*self.ff1)).linear(*self.ff2).dropout(self.dropout)
|
||||
else:
|
||||
x = x + self.attn(x).dropout(self.dropout)
|
||||
x = x.layernorm().linear(*self.ln1)
|
||||
x = x + self.act(x.linear(*self.ff1)).linear(*self.ff2).dropout(self.dropout)
|
||||
x = x.layernorm().linear(*self.ln2)
|
||||
return x
|
||||
|
||||
class Transformer:
|
||||
def __init__(self, syms, maxlen, layers, embed_dim, num_heads, ff_dim):
|
||||
self.maxlen, self.syms = maxlen, syms
|
||||
self.embed = Tensor.scaled_uniform(maxlen+syms, embed_dim).is_param_(False)
|
||||
self.tbs = [TransformerBlock(embed_dim, num_heads, ff_dim) for _ in range(layers)]
|
||||
self.final = Tensor.scaled_uniform(embed_dim, syms)
|
||||
|
||||
def forward(self, x):
|
||||
bs = x.shape[0]
|
||||
|
||||
maxlen_eye = Tensor.eye(x.shape[1])
|
||||
maxlen_eye = maxlen_eye.unsqueeze(0).expand([bs, *maxlen_eye.shape])
|
||||
|
||||
onehot_feat = x.int().one_hot(self.syms)
|
||||
|
||||
onehot = maxlen_eye.cat(onehot_feat, dim=2).flatten(end_dim=1)
|
||||
|
||||
x = onehot.dot(self.embed).reshape((bs, x.shape[1], -1))
|
||||
x = x.sequential(self.tbs)
|
||||
x = x.reshape((-1, x.shape[-1])).dot(self.final).log_softmax()
|
||||
return x.reshape((bs, -1, x.shape[-1]))
|
||||
262
tinygrad_repo/extra/models/unet.py
Normal file
262
tinygrad_repo/extra/models/unet.py
Normal file
@@ -0,0 +1,262 @@
|
||||
from tinygrad import Tensor, Device, dtypes, nn
|
||||
from typing import Optional, Union, List, Any, Tuple, Callable
|
||||
import math
|
||||
|
||||
# allow for monkeypatching
|
||||
Linear, Conv2d, GroupNorm, LayerNorm = nn.Linear, nn.Conv2d, nn.GroupNorm, nn.LayerNorm
|
||||
attention, gelu, mixed_precision_dtype = Tensor.scaled_dot_product_attention, Tensor.gelu, dtypes.float16
|
||||
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/diffusionmodules/util.py#L207
|
||||
def timestep_embedding(timesteps:Tensor, dim:int, max_period=10000):
|
||||
half = dim // 2
|
||||
freqs = (-math.log(max_period) * Tensor.arange(half, device=timesteps.device) / half).exp()
|
||||
args = timesteps.unsqueeze(1) * freqs.unsqueeze(0)
|
||||
out = Tensor.cat(args.cos(), args.sin(), dim=-1)
|
||||
return out.cast(mixed_precision_dtype) if mixed_precision_dtype in Device[Device.DEFAULT].renderer.supported_dtypes() else out
|
||||
|
||||
class ResBlock:
|
||||
def __init__(self, channels:int, emb_channels:int, out_channels:int, num_groups:int=32):
|
||||
self.in_layers = [
|
||||
GroupNorm(num_groups, channels),
|
||||
Tensor.silu,
|
||||
Conv2d(channels, out_channels, 3, padding=1),
|
||||
]
|
||||
self.emb_layers = [
|
||||
Tensor.silu,
|
||||
Linear(emb_channels, out_channels),
|
||||
]
|
||||
self.out_layers = [
|
||||
GroupNorm(num_groups, out_channels),
|
||||
Tensor.silu,
|
||||
lambda x: x, # needed for weights loading code to work
|
||||
Conv2d(out_channels, out_channels, 3, padding=1),
|
||||
]
|
||||
self.skip_connection = Conv2d(channels, out_channels, 1) if channels != out_channels else (lambda x: x)
|
||||
|
||||
def __call__(self, x:Tensor, emb:Tensor) -> Tensor:
|
||||
h = x.sequential(self.in_layers)
|
||||
emb_out = emb.sequential(self.emb_layers)
|
||||
h = h + emb_out.reshape(*emb_out.shape, 1, 1)
|
||||
h = h.sequential(self.out_layers)
|
||||
return self.skip_connection(x) + h
|
||||
|
||||
class CrossAttention:
|
||||
def __init__(self, query_dim:int, ctx_dim:int, n_heads:int, d_head:int):
|
||||
self.to_q = Linear(query_dim, n_heads*d_head, bias=False)
|
||||
self.to_k = Linear(ctx_dim, n_heads*d_head, bias=False)
|
||||
self.to_v = Linear(ctx_dim, n_heads*d_head, bias=False)
|
||||
self.num_heads = n_heads
|
||||
self.head_size = d_head
|
||||
self.attn = attention
|
||||
self.to_out = [Linear(n_heads*d_head, query_dim)]
|
||||
|
||||
def __call__(self, x:Tensor, ctx:Optional[Tensor]=None) -> Tensor:
|
||||
ctx = x if ctx is None else ctx
|
||||
q,k,v = self.to_q(x), self.to_k(ctx), self.to_v(ctx)
|
||||
q,k,v = [y.reshape(x.shape[0], -1, self.num_heads, self.head_size).transpose(1,2) for y in (q,k,v)]
|
||||
attention = self.attn(q, k, v).transpose(1,2)
|
||||
h_ = attention.reshape(x.shape[0], -1, self.num_heads * self.head_size)
|
||||
return h_.sequential(self.to_out)
|
||||
|
||||
class GEGLU:
|
||||
def __init__(self, dim_in:int, dim_out:int):
|
||||
self.proj = Linear(dim_in, dim_out * 2)
|
||||
self.gelu = gelu
|
||||
self.dim_out = dim_out
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
x, gate = self.proj(x).chunk(2, dim=-1)
|
||||
return x * self.gelu(gate)
|
||||
|
||||
class FeedForward:
|
||||
def __init__(self, dim:int, mult:int=4):
|
||||
self.net: tuple[GEGLU, Callable, nn.Linear] = (
|
||||
GEGLU(dim, dim*mult),
|
||||
lambda x: x, # needed for weights loading code to work
|
||||
Linear(dim*mult, dim)
|
||||
)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return x.sequential(list(self.net))
|
||||
|
||||
class BasicTransformerBlock:
|
||||
def __init__(self, dim:int, ctx_dim:int, n_heads:int, d_head:int):
|
||||
self.attn1 = CrossAttention(dim, dim, n_heads, d_head)
|
||||
self.ff = FeedForward(dim)
|
||||
self.attn2 = CrossAttention(dim, ctx_dim, n_heads, d_head)
|
||||
self.norm1 = LayerNorm(dim)
|
||||
self.norm2 = LayerNorm(dim)
|
||||
self.norm3 = LayerNorm(dim)
|
||||
|
||||
def __call__(self, x:Tensor, ctx:Optional[Tensor]=None) -> Tensor:
|
||||
x = x + self.attn1(self.norm1(x))
|
||||
x = x + self.attn2(self.norm2(x), ctx=ctx)
|
||||
x = x + self.ff(self.norm3(x))
|
||||
return x
|
||||
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/attention.py#L619
|
||||
class SpatialTransformer:
|
||||
def __init__(self, channels:int, n_heads:int, d_head:int, ctx_dim:Union[int,List[int]], use_linear:bool, depth:int=1,
|
||||
norm_eps:float=1e-5):
|
||||
if isinstance(ctx_dim, int):
|
||||
ctx_dim = [ctx_dim]*depth
|
||||
else:
|
||||
assert isinstance(ctx_dim, list) and depth == len(ctx_dim)
|
||||
self.norm = GroupNorm(32, channels, eps=norm_eps)
|
||||
assert channels == n_heads * d_head
|
||||
self.proj_in = Linear(channels, channels) if use_linear else Conv2d(channels, channels, 1)
|
||||
self.transformer_blocks = [BasicTransformerBlock(channels, ctx_dim[d], n_heads, d_head) for d in range(depth)]
|
||||
self.proj_out = Linear(channels, channels) if use_linear else Conv2d(channels, channels, 1)
|
||||
self.use_linear = use_linear
|
||||
|
||||
def __call__(self, x:Tensor, ctx:Optional[Tensor]=None) -> Tensor:
|
||||
b, c, h, w = x.shape
|
||||
x_in = x
|
||||
x = self.norm(x)
|
||||
ops = [ (lambda z: z.reshape(b, c, h*w).permute(0,2,1)), (lambda z: self.proj_in(z)) ]
|
||||
x = x.sequential(ops if self.use_linear else ops[::-1])
|
||||
for block in self.transformer_blocks:
|
||||
x = block(x, ctx=ctx)
|
||||
ops = [ (lambda z: self.proj_out(z)), (lambda z: z.permute(0,2,1).reshape(b, c, h, w)) ]
|
||||
x = x.sequential(ops if self.use_linear else ops[::-1])
|
||||
return x + x_in
|
||||
|
||||
class Downsample:
|
||||
def __init__(self, channels:int):
|
||||
self.op = Conv2d(channels, channels, 3, stride=2, padding=1)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
return self.op(x)
|
||||
|
||||
class Upsample:
|
||||
def __init__(self, channels:int):
|
||||
self.conv = Conv2d(channels, channels, 3, padding=1)
|
||||
|
||||
def __call__(self, x:Tensor) -> Tensor:
|
||||
bs,c,py,px = x.shape
|
||||
z = x.reshape(bs, c, py, 1, px, 1).expand(bs, c, py, 2, px, 2).reshape(bs, c, py*2, px*2)
|
||||
return self.conv(z)
|
||||
|
||||
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/diffusionmodules/openaimodel.py#L472
|
||||
class UNetModel:
|
||||
def __init__(self, adm_in_ch:Optional[int], in_ch:int, out_ch:int, model_ch:int, attention_resolutions:List[int], num_res_blocks:int,
|
||||
channel_mult:List[int], transformer_depth:List[int], ctx_dim:Union[int,List[int]], use_linear:bool=False, d_head:Optional[int]=None,
|
||||
n_heads:Optional[int]=None, num_groups:int=32, st_norm_eps:float=1e-5):
|
||||
self.model_ch = model_ch
|
||||
self.num_res_blocks = [num_res_blocks] * len(channel_mult)
|
||||
|
||||
self.attention_resolutions = attention_resolutions
|
||||
self.d_head = d_head
|
||||
self.n_heads = n_heads
|
||||
def get_d_and_n_heads(dims:int) -> Tuple[int,int]:
|
||||
if self.d_head is None:
|
||||
assert self.n_heads is not None, f"d_head and n_heads cannot both be None"
|
||||
return dims // self.n_heads, self.n_heads
|
||||
else:
|
||||
assert self.n_heads is None, f"d_head and n_heads cannot both be non-None"
|
||||
return self.d_head, dims // self.d_head
|
||||
|
||||
time_embed_dim = model_ch * 4
|
||||
self.time_embed = [
|
||||
Linear(model_ch, time_embed_dim),
|
||||
Tensor.silu,
|
||||
Linear(time_embed_dim, time_embed_dim),
|
||||
]
|
||||
|
||||
if adm_in_ch is not None:
|
||||
self.label_emb = [
|
||||
[
|
||||
Linear(adm_in_ch, time_embed_dim),
|
||||
Tensor.silu,
|
||||
Linear(time_embed_dim, time_embed_dim),
|
||||
]
|
||||
]
|
||||
|
||||
self.input_blocks: List[Any] = [
|
||||
[Conv2d(in_ch, model_ch, 3, padding=1)]
|
||||
]
|
||||
input_block_channels = [model_ch]
|
||||
ch = model_ch
|
||||
ds = 1
|
||||
for idx, mult in enumerate(channel_mult):
|
||||
for _ in range(self.num_res_blocks[idx]):
|
||||
layers: List[Any] = [
|
||||
ResBlock(ch, time_embed_dim, model_ch*mult, num_groups),
|
||||
]
|
||||
ch = mult * model_ch
|
||||
if ds in attention_resolutions:
|
||||
d_head, n_heads = get_d_and_n_heads(ch)
|
||||
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx], norm_eps=st_norm_eps))
|
||||
|
||||
self.input_blocks.append(layers)
|
||||
input_block_channels.append(ch)
|
||||
|
||||
if idx != len(channel_mult) - 1:
|
||||
self.input_blocks.append([
|
||||
Downsample(ch),
|
||||
])
|
||||
input_block_channels.append(ch)
|
||||
ds *= 2
|
||||
|
||||
d_head, n_heads = get_d_and_n_heads(ch)
|
||||
self.middle_block: List = [
|
||||
ResBlock(ch, time_embed_dim, ch, num_groups),
|
||||
SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[-1], norm_eps=st_norm_eps),
|
||||
ResBlock(ch, time_embed_dim, ch, num_groups),
|
||||
]
|
||||
|
||||
self.output_blocks = []
|
||||
for idx, mult in list(enumerate(channel_mult))[::-1]:
|
||||
for i in range(self.num_res_blocks[idx] + 1):
|
||||
ich = input_block_channels.pop()
|
||||
layers = [
|
||||
ResBlock(ch + ich, time_embed_dim, model_ch*mult, num_groups),
|
||||
]
|
||||
ch = model_ch * mult
|
||||
|
||||
if ds in attention_resolutions:
|
||||
d_head, n_heads = get_d_and_n_heads(ch)
|
||||
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx], norm_eps=st_norm_eps))
|
||||
|
||||
if idx > 0 and i == self.num_res_blocks[idx]:
|
||||
layers.append(Upsample(ch))
|
||||
ds //= 2
|
||||
self.output_blocks.append(layers)
|
||||
|
||||
self.out = [
|
||||
GroupNorm(num_groups, ch),
|
||||
Tensor.silu,
|
||||
Conv2d(model_ch, out_ch, 3, padding=1),
|
||||
]
|
||||
|
||||
def __call__(self, x:Tensor, tms:Tensor, ctx:Tensor, y:Optional[Tensor]=None) -> Tensor:
|
||||
t_emb = timestep_embedding(tms, self.model_ch)
|
||||
emb = t_emb.sequential(self.time_embed)
|
||||
|
||||
if y is not None:
|
||||
assert y.shape[0] == x.shape[0]
|
||||
emb = emb + y.sequential(self.label_emb[0])
|
||||
|
||||
if mixed_precision_dtype in Device[Device.DEFAULT].renderer.supported_dtypes():
|
||||
emb = emb.cast(mixed_precision_dtype)
|
||||
ctx = ctx.cast(mixed_precision_dtype)
|
||||
x = x .cast(mixed_precision_dtype)
|
||||
|
||||
def run(x:Tensor, bb) -> Tensor:
|
||||
if isinstance(bb, ResBlock): x = bb(x, emb)
|
||||
elif isinstance(bb, SpatialTransformer): x = bb(x, ctx)
|
||||
else: x = bb(x)
|
||||
return x
|
||||
|
||||
saved_inputs = []
|
||||
for b in self.input_blocks:
|
||||
for bb in b:
|
||||
x = run(x, bb)
|
||||
saved_inputs.append(x)
|
||||
for bb in self.middle_block:
|
||||
x = run(x, bb)
|
||||
for b in self.output_blocks:
|
||||
x = x.cat(saved_inputs.pop(), dim=1)
|
||||
for bb in b:
|
||||
x = run(x, bb)
|
||||
return x.sequential(self.out)
|
||||
59
tinygrad_repo/extra/models/unet3d.py
Normal file
59
tinygrad_repo/extra/models/unet3d.py
Normal file
@@ -0,0 +1,59 @@
|
||||
from pathlib import Path
|
||||
import torch
|
||||
from tinygrad import nn
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.helpers import fetch, get_child
|
||||
|
||||
class DownsampleBlock:
|
||||
def __init__(self, c0, c1, stride=2):
|
||||
self.conv1 = [nn.Conv2d(c0, c1, kernel_size=(3,3,3), stride=stride, padding=(1,1,1,1,1,1), bias=False), nn.InstanceNorm(c1), Tensor.relu]
|
||||
self.conv2 = [nn.Conv2d(c1, c1, kernel_size=(3,3,3), padding=(1,1,1,1,1,1), bias=False), nn.InstanceNorm(c1), Tensor.relu]
|
||||
|
||||
def __call__(self, x):
|
||||
return x.sequential(self.conv1).sequential(self.conv2)
|
||||
|
||||
class UpsampleBlock:
|
||||
def __init__(self, c0, c1):
|
||||
self.upsample_conv = [nn.ConvTranspose2d(c0, c1, kernel_size=(2,2,2), stride=2)]
|
||||
self.conv1 = [nn.Conv2d(2 * c1, c1, kernel_size=(3,3,3), padding=(1,1,1,1,1,1), bias=False), nn.InstanceNorm(c1), Tensor.relu]
|
||||
self.conv2 = [nn.Conv2d(c1, c1, kernel_size=(3,3,3), padding=(1,1,1,1,1,1), bias=False), nn.InstanceNorm(c1), Tensor.relu]
|
||||
|
||||
def __call__(self, x, skip):
|
||||
x = x.sequential(self.upsample_conv)
|
||||
x = Tensor.cat(x, skip, dim=1)
|
||||
return x.sequential(self.conv1).sequential(self.conv2)
|
||||
|
||||
class UNet3D:
|
||||
def __init__(self, in_channels=1, n_class=3):
|
||||
filters = [32, 64, 128, 256, 320]
|
||||
inp, out = filters[:-1], filters[1:]
|
||||
self.input_block = DownsampleBlock(in_channels, filters[0], stride=1)
|
||||
self.downsample = [DownsampleBlock(i, o) for i, o in zip(inp, out)]
|
||||
self.bottleneck = DownsampleBlock(filters[-1], filters[-1])
|
||||
self.upsample = [UpsampleBlock(filters[-1], filters[-1])] + [UpsampleBlock(i, o) for i, o in zip(out[::-1], inp[::-1])]
|
||||
self.output = {"conv": nn.Conv2d(filters[0], n_class, kernel_size=(1, 1, 1))}
|
||||
|
||||
def __call__(self, x):
|
||||
x = self.input_block(x)
|
||||
outputs = [x]
|
||||
for downsample in self.downsample:
|
||||
x = downsample(x)
|
||||
outputs.append(x)
|
||||
x = self.bottleneck(x)
|
||||
for upsample, skip in zip(self.upsample, outputs[::-1]):
|
||||
x = upsample(x, skip)
|
||||
x = self.output["conv"](x)
|
||||
return x
|
||||
|
||||
def load_from_pretrained(self):
|
||||
fn = Path(__file__).parents[1] / "weights" / "unet-3d.ckpt"
|
||||
fetch("https://zenodo.org/record/5597155/files/3dunet_kits19_pytorch.ptc?download=1", fn)
|
||||
state_dict = torch.jit.load(fn, map_location=torch.device("cpu")).state_dict()
|
||||
for k, v in state_dict.items():
|
||||
obj = get_child(self, k)
|
||||
assert obj.shape == v.shape, (k, obj.shape, v.shape)
|
||||
obj.assign(v.numpy())
|
||||
|
||||
if __name__ == "__main__":
|
||||
mdl = UNet3D()
|
||||
mdl.load_from_pretrained()
|
||||
73
tinygrad_repo/extra/models/vit.py
Normal file
73
tinygrad_repo/extra/models/vit.py
Normal file
@@ -0,0 +1,73 @@
|
||||
import numpy as np
|
||||
from tinygrad.tensor import Tensor
|
||||
from tinygrad.helpers import fetch
|
||||
from extra.models.transformer import TransformerBlock
|
||||
|
||||
class ViT:
|
||||
def __init__(self, layers=12, embed_dim=192, num_heads=3):
|
||||
self.embedding = (Tensor.uniform(embed_dim, 3, 16, 16), Tensor.zeros(embed_dim))
|
||||
self.embed_dim = embed_dim
|
||||
self.cls = Tensor.ones(1, 1, embed_dim)
|
||||
self.pos_embedding = Tensor.ones(1, 197, embed_dim)
|
||||
self.tbs = [
|
||||
TransformerBlock(embed_dim=embed_dim, num_heads=num_heads, ff_dim=embed_dim*4,
|
||||
prenorm=True, act=lambda x: x.gelu())
|
||||
for i in range(layers)]
|
||||
self.encoder_norm = (Tensor.uniform(embed_dim), Tensor.zeros(embed_dim))
|
||||
self.head = (Tensor.uniform(embed_dim, 1000), Tensor.zeros(1000))
|
||||
|
||||
def patch_embed(self, x):
|
||||
x = x.conv2d(*self.embedding, stride=16)
|
||||
x = x.reshape(shape=(x.shape[0], x.shape[1], -1)).permute(order=(0,2,1))
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
ce = self.cls.add(Tensor.zeros(x.shape[0],1,1))
|
||||
pe = self.patch_embed(x)
|
||||
x = ce.cat(pe, dim=1)
|
||||
x = x.add(self.pos_embedding).sequential(self.tbs)
|
||||
x = x.layernorm().linear(*self.encoder_norm)
|
||||
return x[:, 0].linear(*self.head)
|
||||
|
||||
def load_from_pretrained(m):
|
||||
# https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/vision_transformer.py
|
||||
if m.embed_dim == 192:
|
||||
url = "https://storage.googleapis.com/vit_models/augreg/Ti_16-i21k-300ep-lr_0.001-aug_none-wd_0.03-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.03-res_224.npz"
|
||||
elif m.embed_dim == 768:
|
||||
url = "https://storage.googleapis.com/vit_models/augreg/B_16-i21k-300ep-lr_0.001-aug_medium1-wd_0.1-do_0.0-sd_0.0--imagenet2012-steps_20k-lr_0.01-res_224.npz"
|
||||
else:
|
||||
raise Exception("no pretrained weights for configuration")
|
||||
dat = np.load(fetch(url))
|
||||
|
||||
#for x in dat.keys():
|
||||
# print(x, dat[x].shape, dat[x].dtype)
|
||||
|
||||
m.embedding[0].assign(np.transpose(dat['embedding/kernel'], (3,2,0,1)))
|
||||
m.embedding[1].assign(dat['embedding/bias'])
|
||||
|
||||
m.cls.assign(dat['cls'])
|
||||
|
||||
m.head[0].assign(dat['head/kernel'])
|
||||
m.head[1].assign(dat['head/bias'])
|
||||
|
||||
m.pos_embedding.assign(dat['Transformer/posembed_input/pos_embedding'])
|
||||
m.encoder_norm[0].assign(dat['Transformer/encoder_norm/scale'])
|
||||
m.encoder_norm[1].assign(dat['Transformer/encoder_norm/bias'])
|
||||
|
||||
for i in range(12):
|
||||
m.tbs[i].query[0].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/query/kernel'].reshape(m.embed_dim, m.embed_dim))
|
||||
m.tbs[i].query[1].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/query/bias'].reshape(m.embed_dim))
|
||||
m.tbs[i].key[0].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/key/kernel'].reshape(m.embed_dim, m.embed_dim))
|
||||
m.tbs[i].key[1].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/key/bias'].reshape(m.embed_dim))
|
||||
m.tbs[i].value[0].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/value/kernel'].reshape(m.embed_dim, m.embed_dim))
|
||||
m.tbs[i].value[1].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/value/bias'].reshape(m.embed_dim))
|
||||
m.tbs[i].out[0].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/out/kernel'].reshape(m.embed_dim, m.embed_dim))
|
||||
m.tbs[i].out[1].assign(dat[f'Transformer/encoderblock_{i}/MultiHeadDotProductAttention_1/out/bias'].reshape(m.embed_dim))
|
||||
m.tbs[i].ff1[0].assign(dat[f'Transformer/encoderblock_{i}/MlpBlock_3/Dense_0/kernel'])
|
||||
m.tbs[i].ff1[1].assign(dat[f'Transformer/encoderblock_{i}/MlpBlock_3/Dense_0/bias'])
|
||||
m.tbs[i].ff2[0].assign(dat[f'Transformer/encoderblock_{i}/MlpBlock_3/Dense_1/kernel'])
|
||||
m.tbs[i].ff2[1].assign(dat[f'Transformer/encoderblock_{i}/MlpBlock_3/Dense_1/bias'])
|
||||
m.tbs[i].ln1[0].assign(dat[f'Transformer/encoderblock_{i}/LayerNorm_0/scale'])
|
||||
m.tbs[i].ln1[1].assign(dat[f'Transformer/encoderblock_{i}/LayerNorm_0/bias'])
|
||||
m.tbs[i].ln2[0].assign(dat[f'Transformer/encoderblock_{i}/LayerNorm_2/scale'])
|
||||
m.tbs[i].ln2[1].assign(dat[f'Transformer/encoderblock_{i}/LayerNorm_2/bias'])
|
||||
Reference in New Issue
Block a user