vedaco commited on
Commit
e2e6e07
·
verified ·
1 Parent(s): 5a827a2

Create akasha/generate.py

Browse files
Files changed (1) hide show
  1. akasha/generate.py +109 -0
akasha/generate.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Autoregressive generation for AKASHA.
3
+ Generates images token-by-token using the transformer.
4
+ """
5
+
6
+ import tensorflow as tf
7
+ import numpy as np
8
+ from tqdm import tqdm
9
+
10
+
11
+ class AKASHAGenerator:
12
+ """Handles autoregressive image generation."""
13
+
14
+ def __init__(self, model, config):
15
+ self.model = model
16
+ self.config = config
17
+ tok = config["model"]["tokenizer"]
18
+ self.image_size = tok["image_size"]
19
+ self.patch_size = tok["patch_size"]
20
+ self.num_tokens = tok["num_tokens"]
21
+ self.grid_size = self.image_size // self.patch_size
22
+ self.seq_length = self.grid_size * self.grid_size
23
+ self.bos_token = self.num_tokens
24
+ self.eos_token = self.num_tokens + 1
25
+
26
+ def generate(self, num_images=1, temperature=0.9, top_k=100,
27
+ top_p=0.95, seed=None, callback=None):
28
+ """Generate images autoregressively."""
29
+ if seed is not None:
30
+ tf.random.set_seed(seed)
31
+
32
+ current_tokens = tf.cast(
33
+ tf.fill([num_images, 1], self.bos_token), tf.int32
34
+ )
35
+
36
+ generated_tokens = []
37
+ desc = f"Generating {num_images} image(s) ({self.seq_length} tokens)"
38
+
39
+ for step in tqdm(range(self.seq_length), desc=desc):
40
+ logits = self.model.transformer(current_tokens, training=False)
41
+ next_token_logits = logits[:, -1, :] / temperature
42
+
43
+ # Top-k filtering
44
+ if top_k > 0:
45
+ top_k_actual = min(top_k, next_token_logits.shape[-1])
46
+ values, _ = tf.math.top_k(next_token_logits, k=top_k_actual)
47
+ min_value = values[:, -1:]
48
+ next_token_logits = tf.where(
49
+ next_token_logits < min_value,
50
+ tf.fill(tf.shape(next_token_logits), -1e10),
51
+ next_token_logits,
52
+ )
53
+
54
+ next_token = tf.random.categorical(
55
+ next_token_logits, num_samples=1, dtype=tf.int32
56
+ )
57
+ next_token = tf.clip_by_value(next_token, 0, self.num_tokens - 1)
58
+
59
+ generated_tokens.append(next_token)
60
+ current_tokens = tf.concat([current_tokens, next_token], axis=1)
61
+
62
+ if callback:
63
+ callback(step + 1, self.seq_length)
64
+
65
+ token_sequence = tf.concat(generated_tokens, axis=1)
66
+
67
+ from akasha.tokenizer import ImageTokenizer
68
+ tokenizer = ImageTokenizer(self.model.vqvae)
69
+ images = tokenizer.detokenize(token_sequence, grid_size=self.grid_size)
70
+
71
+ return images, token_sequence
72
+
73
+ def generate_with_kv_cache(self, num_images=1, temperature=0.9,
74
+ top_k=100, top_p=0.95):
75
+ """Generate with KV cache for faster inference."""
76
+ current_token = tf.cast(
77
+ tf.fill([num_images, 1], self.bos_token), tf.int32
78
+ )
79
+ past_kvs = [None] * self.model.transformer.num_layers
80
+ generated_tokens = []
81
+
82
+ for step in tqdm(range(self.seq_length), desc="Generating (cached)"):
83
+ logits, past_kvs = self.model.transformer(
84
+ current_token, training=False, use_cache=True, past_kvs=past_kvs
85
+ )
86
+ next_token_logits = logits[:, -1, :] / temperature
87
+
88
+ if top_k > 0:
89
+ top_k_actual = min(top_k, next_token_logits.shape[-1])
90
+ values, _ = tf.math.top_k(next_token_logits, k=top_k_actual)
91
+ min_value = values[:, -1:]
92
+ next_token_logits = tf.where(
93
+ next_token_logits < min_value,
94
+ tf.fill(tf.shape(next_token_logits), -1e10),
95
+ next_token_logits,
96
+ )
97
+
98
+ next_token = tf.random.categorical(
99
+ next_token_logits, num_samples=1, dtype=tf.int32
100
+ )
101
+ next_token = tf.clip_by_value(next_token, 0, self.num_tokens - 1)
102
+ generated_tokens.append(next_token)
103
+ current_token = next_token
104
+
105
+ token_sequence = tf.concat(generated_tokens, axis=1)
106
+ from akasha.tokenizer import ImageTokenizer
107
+ tokenizer = ImageTokenizer(self.model.vqvae)
108
+ images = tokenizer.detokenize(token_sequence, grid_size=self.grid_size)
109
+ return images, token_sequence