Andhs commited on
Commit
dc3806d
·
verified ·
1 Parent(s): 9ebf833

Upload app.py

Browse files
Files changed (1) hide show
  1. app.py +522 -467
app.py CHANGED
@@ -1,467 +1,522 @@
1
- import sys
2
- import types
3
- import os
4
- import math
5
- import copy
6
- import json
7
- import torch
8
- import torch.nn.functional as F
9
- from flask import Flask, request, jsonify, Response
10
- from transformers import AutoTokenizer, AutoModelForMaskedLM, AutoModelForCausalLM
11
-
12
- # 1. Environment Parsing & Architecture Strategy Mapping
13
- MODEL_NAME = os.getenv("MODEL_NAME", "dllm-hub/Qwen3-0.6B-diffusion-bd3lm-v0.1")
14
- IS_DIFFUSION = "diffusion" in MODEL_NAME.lower()
15
-
16
- # Dynamic initialization layer targeting Diffusion Language Models
17
- if IS_DIFFUSION:
18
- try:
19
- import dllm.utils
20
- import dllm.pipelines
21
- import dllm.data
22
- import dllm.core
23
- except ImportError:
24
- pass
25
- if 'dllm' not in sys.modules:
26
- dllm_mock = types.ModuleType('dllm')
27
- dllm_mock.core = sys.modules.get('dllm.core')
28
- dllm_mock.data = sys.modules.get('dllm.data')
29
- dllm_mock.pipelines = sys.modules.get('dllm.pipelines')
30
- dllm_mock.utils = sys.modules.get('dllm.utils')
31
- sys.modules['dllm'] = dllm_mock
32
-
33
- app = Flask(__name__)
34
- model = None
35
- tokenizer = None
36
- device = None
37
-
38
- # ==========================================================
39
- # SYSTEM WORKSPACE PIPELINES: CORE DIFFUSION SAMPLING LOOPS
40
- # ==========================================================
41
- def add_gumbel_noise(logits, temperature):
42
- if temperature == 0:
43
- return logits
44
- logits = logits.to(torch.float64)
45
- noise = torch.rand_like(logits, dtype=torch.float64)
46
- g = (-torch.log(noise)) ** temperature
47
- return logits.exp() / g
48
-
49
- def get_num_transfer_tokens(mask_index, steps):
50
- mask_num = mask_index.sum(dim=1, keepdim=True)
51
- base = mask_num // steps
52
- rem = mask_num % steps
53
- out = torch.zeros(mask_num.size(0), steps, device=mask_index.device, dtype=torch.long) + base
54
- for i in range(mask_num.size(0)):
55
- out[i, : rem[i]] += 1
56
- return out
57
-
58
- def build_staircase_attention_mask(x, block_size, pad_id):
59
- B, T = x.shape
60
- device = x.device
61
- valid = x != pad_id
62
- pos_raw = torch.cumsum(valid.long(), dim=-1)
63
- position_ids = torch.where(valid, pos_raw - 1, torch.zeros_like(pos_raw)).long()
64
- col = torch.arange(T, device=device)
65
- block_ids = (col // block_size).view(1, T).expand(B, T)
66
- block_ids = torch.where(valid, block_ids, torch.full_like(block_ids, -1))
67
- q = block_ids.view(B, 1, T, 1)
68
- k = block_ids.view(B, 1, 1, T)
69
- attn = (k <= q) & (q >= 0) & (k >= 0)
70
- return attn, position_ids
71
-
72
- def diffusion_step_block(logits, x_block, mask_block, num_transfer, temperature, remasking):
73
- B, L, _ = logits.shape
74
- if not mask_block.any():
75
- return x_block
76
- noisy = add_gumbel_noise(logits, temperature)
77
- x0 = noisy.argmax(dim=-1)
78
- if remasking == "low_confidence":
79
- p = F.softmax(logits, dim=-1)
80
- conf = p.gather(-1, x0.unsqueeze(-1)).squeeze(-1)
81
- elif remasking == "random":
82
- conf = torch.rand((B, L), device=logits.device)
83
- else:
84
- raise ValueError(remasking)
85
- x0 = torch.where(mask_block, x0, x_block)
86
- neg_inf = torch.full_like(conf, -float("inf"))
87
- conf = torch.where(mask_block, conf, neg_inf)
88
- commit = torch.zeros_like(x_block, dtype=torch.bool)
89
- for i in range(B):
90
- k = int(num_transfer[i].item())
91
- if k > 0:
92
- valid = (conf[i] > -float("inf")).sum().item()
93
- k = min(k, valid)
94
- _, idx = torch.topk(conf[i], k)
95
- commit[i, idx] = True
96
- out = x_block.clone()
97
- out[commit] = x0[commit]
98
- return out
99
-
100
- @torch.no_grad()
101
- def generate(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size=32, temperature=0.0, cfg_scale=0.0, remasking="low_confidence", capture_interval=0):
102
- device = model.device
103
- mask_id = tokenizer.mask_token_id
104
- pad_id = tokenizer.pad_token_id
105
- if pad_id is None:
106
- pad_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else tokenizer.mask_token_id
107
- if isinstance(prompt, torch.Tensor):
108
- x = prompt.to(device).long()
109
- else:
110
- if isinstance(prompt[0], (list, tuple)):
111
- max_len = max(len(p) for p in prompt)
112
- x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
113
- for i, p in enumerate(prompt):
114
- x[i, : len(p)] = torch.tensor(p, device=device)
115
- else:
116
- x = torch.tensor(prompt, device=device).long()
117
- if x.dim() == 1:
118
- x = x.unsqueeze(0)
119
- B = x.size(0)
120
- finished = torch.zeros(B, dtype=torch.bool, device=device)
121
- num_blocks = math.ceil(max_new_tokens / block_size)
122
- steps_per_block = math.ceil(steps / num_blocks)
123
- generated = 0
124
- intermediates = []
125
- total_step = 0
126
- while generated < max_new_tokens:
127
- if finished.all():
128
- break
129
- T_prefix = x.size(1)
130
- offset = T_prefix % block_size
131
- room = block_size if offset == 0 else block_size - offset
132
- cur_len = min(room, max_new_tokens - generated)
133
- if cur_len <= 0:
134
- break
135
- attn_pfx, pos_pfx = build_staircase_attention_mask(x, block_size, pad_id)
136
- out = model(x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
137
- cond_past = out.past_key_values
138
- if cfg_scale > 0:
139
- un_x = x.clone()
140
- un_x[:] = mask_id
141
- out_un = model(un_x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
142
- uncond_past = out_un.past_key_values
143
- else:
144
- uncond_past = None
145
- block = torch.full((B, cur_len), mask_id, device=device, dtype=torch.long)
146
- block[finished] = pad_id
147
- x = torch.cat([x, block], dim=1)
148
- T_total = x.size(1)
149
- block_mask = x[:, -cur_len:] == mask_id
150
- num_transfer = get_num_transfer_tokens(block_mask, steps_per_block)
151
- eff_steps = num_transfer.size(1)
152
- full_attn, full_pos = build_staircase_attention_mask(x, block_size, pad_id)
153
- attn_blk = full_attn[:, :, T_prefix:T_total, :]
154
- pos_blk = full_pos[:, T_prefix:T_total]
155
- for t in range(eff_steps):
156
- x_blk = x[:, T_prefix:T_total]
157
- m_blk = x_blk == mask_id
158
- cond_logits = model(x_blk, attention_mask=attn_blk, position_ids=pos_blk, past_key_values=copy.deepcopy(cond_past), use_cache=False).logits
159
- logits = cond_logits
160
- if cfg_scale > 0:
161
- un_logits = model(x_blk, attention_mask=attn_blk, position_ids=pos_blk, past_key_values=copy.deepcopy(uncond_past), use_cache=False).logits
162
- logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
163
- x_blk_new = diffusion_step_block(logits, x_blk, m_blk, num_transfer[:, t], temperature, remasking)
164
- x[:, T_prefix:T_total] = x_blk_new
165
- if capture_interval > 0 and total_step % capture_interval == 0:
166
- intermediates.append(x.clone())
167
- total_step += 1
168
- if tokenizer.eos_token_id is not None:
169
- finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
170
- generated += cur_len
171
- if finished.all():
172
- break
173
-
174
- if capture_interval > 0:
175
- return x, intermediates
176
- return x
177
-
178
- @torch.no_grad()
179
- def generate_stream(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size=32, temperature=0.0, cfg_scale=0.0, remasking="low_confidence", capture_interval=10):
180
- device = model.device
181
- mask_id = tokenizer.mask_token_id
182
- pad_id = tokenizer.pad_token_id
183
- if pad_id is None:
184
- pad_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else tokenizer.mask_token_id
185
- if isinstance(prompt, torch.Tensor):
186
- x = prompt.to(device).long()
187
- else:
188
- if isinstance(prompt[0], (list, tuple)):
189
- max_len = max(len(p) for p in prompt)
190
- x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
191
- for i, p in enumerate(prompt):
192
- x[i, : len(p)] = torch.tensor(p, device=device)
193
- else:
194
- x = torch.tensor(prompt, device=device).long()
195
- if x.dim() == 1:
196
- x = x.unsqueeze(0)
197
- B = x.size(0)
198
- finished = torch.zeros(B, dtype=torch.bool, device=device)
199
- num_blocks = math.ceil(max_new_tokens / block_size)
200
- steps_per_block = math.ceil(steps / num_blocks)
201
- generated = 0
202
- total_step = 0
203
- prompt_len = x.size(1)
204
- while generated < max_new_tokens:
205
- if finished.all():
206
- break
207
- T_prefix = x.size(1)
208
- offset = T_prefix % block_size
209
- room = block_size if offset == 0 else block_size - offset
210
- cur_len = min(room, max_new_tokens - generated)
211
- if cur_len <= 0:
212
- break
213
- attn_pfx, pos_pfx = build_staircase_attention_mask(x, block_size, pad_id)
214
- out = model(x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
215
- cond_past = out.past_key_values
216
- if cfg_scale > 0:
217
- un_x = x.clone()
218
- un_x[:] = mask_id
219
- out_un = model(un_x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
220
- uncond_past = out_un.past_key_values
221
- else:
222
- uncond_past = None
223
- block = torch.full((B, cur_len), mask_id, device=device, dtype=torch.long)
224
- block[finished] = pad_id
225
- x = torch.cat([x, block], dim=1)
226
- T_total = x.size(1)
227
- block_mask = x[:, -cur_len:] == mask_id
228
- num_transfer = get_num_transfer_tokens(block_mask, steps_per_block)
229
- eff_steps = num_transfer.size(1)
230
- full_attn, full_pos = build_staircase_attention_mask(x, block_size, pad_id)
231
- attn_blk = full_attn[:, :, T_prefix:T_total, :]
232
- pos_blk = full_pos[:, T_prefix:T_total]
233
- for t in range(eff_steps):
234
- x_blk = x[:, T_prefix:T_total]
235
- m_blk = x_blk == mask_id
236
- cond_logits = model(x_blk, attention_mask=attn_blk, position_ids=pos_blk, past_key_values=copy.deepcopy(cond_past), use_cache=False).logits
237
- logits = cond_logits
238
- if cfg_scale > 0:
239
- un_logits = model(
240
- x_blk, attention_mask=attn_blk, position_ids=pos_blk,
241
- past_key_values=copy.deepcopy(uncond_past), use_cache=False
242
- ).logits
243
- logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
244
- x_blk_new = diffusion_step_block(
245
- logits, x_blk, m_blk, num_transfer[:, t], temperature, remasking
246
- )
247
- x[:, T_prefix:T_total] = x_blk_new
248
-
249
- if total_step % capture_interval == 0:
250
- new_tokens = x[0, prompt_len:prompt_len + max_new_tokens].tolist()
251
- text = tokenizer.decode(new_tokens, skip_special_tokens=True)
252
- yield {
253
- "type": "intermediate",
254
- "step": total_step,
255
- "text": text,
256
- "total_steps": steps
257
- }
258
-
259
- total_step += 1
260
-
261
- if tokenizer.eos_token_id is not None:
262
- finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
263
- if finished.all():
264
- break
265
- generated += cur_len
266
-
267
- if finished.all():
268
- break
269
-
270
- new_tokens = x[0, prompt_len:prompt_len + max_new_tokens].tolist()
271
- final_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
272
- yield {
273
- "type": "final",
274
- "text": final_text,
275
- "total_steps": total_step
276
- }
277
-
278
-
279
- # ==========================================================
280
- # ARCHITECTURE ROUTING LAYERS & TRANSLATION ENGINE CODES
281
- # ==========================================================
282
- def load_model():
283
- global model, tokenizer, device
284
- device = "cuda" if torch.cuda.is_available() else "cpu"
285
-
286
- print(f"Initializing {MODEL_NAME} on {device}... (Diffusion Strategy Flag = {IS_DIFFUSION})")
287
-
288
- if IS_DIFFUSION:
289
- model = AutoModelForMaskedLM.from_pretrained(
290
- MODEL_NAME,
291
- dtype=torch.bfloat16,
292
- trust_remote_code=True
293
- ).to(device).eval()
294
- else:
295
- model = AutoModelForCausalLM.from_pretrained(
296
- MODEL_NAME,
297
- torch_dtype=torch.bfloat16,
298
- trust_remote_code=False
299
- ).to(device).eval()
300
-
301
- tokenizer = AutoTokenizer.from_pretrained(
302
- MODEL_NAME,
303
- trust_remote_code=IS_DIFFUSION
304
- )
305
- print("Model compilation completed and loaded into memory workspace.")
306
-
307
- @app.route('/health', methods=['GET'])
308
- def health():
309
- return jsonify({"status": "healthy", "model_loaded": model is not None, "is_diffusion": IS_DIFFUSION})
310
-
311
- @app.route('/generate', methods=['POST'])
312
- def generate_text():
313
- if model is None or tokenizer is None:
314
- return jsonify({"error": "Model initialization missing"}), 503
315
-
316
- data = request.get_json() or {}
317
- if 'prompt' not in data:
318
- return jsonify({"error": "Missing 'prompt' operational field"}), 400
319
-
320
- prompt = data['prompt']
321
- max_new_tokens = data.get('max_new_tokens', 256)
322
- temperature = data.get('temperature', 0.0)
323
- system_prompt = data.get('system_prompt', 'You are an expert real-time translation assistant.')
324
-
325
- messages = [
326
- {"role": "system", "content": system_prompt},
327
- {"role": "user", "content": prompt}
328
- ]
329
-
330
- encoded = tokenizer.apply_chat_template(
331
- messages,
332
- add_generation_prompt=True,
333
- tokenize=True,
334
- enable_thinking=False if IS_DIFFUSION else None
335
- )
336
-
337
- if IS_DIFFUSION:
338
- input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
339
- steps = data.get('steps', 256)
340
- block_size = data.get('block_size', 32)
341
- cfg_scale = data.get('cfg_scale', 0.0)
342
- remasking = data.get('remasking', 'low_confidence')
343
-
344
- output = generate(
345
- model, tokenizer, input_ids,
346
- steps=steps, max_new_tokens=max_new_tokens, block_size=block_size,
347
- temperature=temperature, cfg_scale=cfg_scale, remasking=remasking,
348
- )
349
- prompt_len = len(encoded)
350
- new_tokens = output[0, prompt_len:prompt_len + max_new_tokens].tolist()
351
- generated_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
352
- else:
353
- # High-Speed Autoregressive Optimization Matrix
354
- input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
355
- with torch.no_grad():
356
- output_ids = model.generate(
357
- input_ids,
358
- max_new_tokens=max_new_tokens,
359
- temperature=temperature,
360
- do_sample=True if temperature > 0 else False,
361
- pad_token_id=tokenizer.eos_token_id
362
- )
363
- generated_ids = output_ids[0, input_ids.shape[-1]:]
364
- generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
365
-
366
- return jsonify({"prompt": prompt, "generated_text": generated_text})
367
-
368
- @app.route('/generate_stream', methods=['POST'])
369
- def generate_text_stream():
370
- if model is None or tokenizer is None:
371
- return jsonify({"error": "Model workspace offline"}), 503
372
-
373
- data = request.get_json() or {}
374
- if not data or 'prompt' not in data:
375
- return jsonify({"error": "Missing 'prompt' operational field"}), 400
376
-
377
- prompt = data['prompt']
378
- max_new_tokens = data.get('max_new_tokens', 256)
379
- temperature = data.get('temperature', 0.0)
380
- system_prompt = data.get('system_prompt', 'You are an expert real-time translation assistant.')
381
-
382
- messages = [
383
- {"role": "system", "content": system_prompt},
384
- {"role": "user", "content": prompt}
385
- ]
386
- encoded = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=True, enable_thinking=False if IS_DIFFUSION else None)
387
-
388
- if IS_DIFFUSION:
389
- input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
390
- steps = data.get('steps', 256)
391
- block_size = data.get('block_size', 32)
392
- cfg_scale = data.get('cfg_scale', 0.0)
393
- remasking = data.get('remasking', 'low_confidence')
394
- capture_interval = data.get('capture_interval', 10)
395
-
396
- output, intermediates = generate(
397
- model, tokenizer, input_ids,
398
- steps=steps, max_new_tokens=max_new_tokens, block_size=block_size,
399
- temperature=temperature, cfg_scale=cfg_scale, remasking=remasking,
400
- capture_interval=capture_interval,
401
- )
402
- prompt_len = len(encoded)
403
- intermediate_states = []
404
- for i, intermediate in enumerate(intermediates):
405
- new_tokens = intermediate[0, prompt_len:prompt_len + max_new_tokens].tolist()
406
- text = tokenizer.decode(new_tokens, skip_special_tokens=True)
407
- intermediate_states.append({"step": i * capture_interval, "text": text})
408
-
409
- new_tokens = output[0, prompt_len:prompt_len + max_new_tokens].tolist()
410
- generated_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
411
- return jsonify({"prompt": prompt, "generated_text": generated_text, "intermediate_states": intermediate_states})
412
- else:
413
- input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
414
- with torch.no_grad():
415
- output_ids = model.generate(input_ids, max_new_tokens=max_new_tokens, temperature=temperature, do_sample=True if temperature > 0 else False, pad_token_id=tokenizer.eos_token_id)
416
- generated_ids = output_ids[0, input_ids.shape[-1]:]
417
- generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
418
- return jsonify({"prompt": prompt, "generated_text": generated_text, "intermediate_states": []})
419
-
420
- @app.route('/generate_sse', methods=['POST'])
421
- def generate_text_sse():
422
- if model is None or tokenizer is None:
423
- return jsonify({"error": "Model workspace offline"}), 503
424
-
425
- data = request.get_json() or {}
426
- if not data or 'prompt' not in data:
427
- return jsonify({"error": "Missing 'prompt' operational field"}), 400
428
-
429
- prompt = data['prompt']
430
- max_new_tokens = data.get('max_new_tokens', 256)
431
- temperature = data.get('temperature', 0.0)
432
- system_prompt = data.get('system_prompt', 'You are an expert real-time translation assistant.')
433
-
434
- messages = [
435
- {"role": "system", "content": system_prompt},
436
- {"role": "user", "content": prompt}
437
- ]
438
- encoded = tokenizer.apply_chat_template(messages, add_generation_prompt=True, tokenize=True, enable_thinking=False if IS_DIFFUSION else None)
439
- input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
440
-
441
- def stream():
442
- if IS_DIFFUSION:
443
- steps = data.get('steps', 256)
444
- block_size = data.get('block_size', 32)
445
- cfg_scale = data.get('cfg_scale', 0.0)
446
- remasking = data.get('remasking', 'low_confidence')
447
- capture_interval = data.get('capture_interval', 10)
448
- for state in generate_stream(model, tokenizer, input_ids, steps=steps, max_new_tokens=max_new_tokens, block_size=block_size, temperature=temperature, cfg_scale=cfg_scale, remasking=remasking, capture_interval=capture_interval):
449
- yield f"data: {json.dumps(state)}\n\n"
450
- else:
451
- with torch.no_grad():
452
- output_ids = model.generate(input_ids, max_new_tokens=max_new_tokens, temperature=temperature, do_sample=True if temperature > 0 else False, pad_token_id=tokenizer.eos_token_id)
453
- generated_ids = output_ids[0, input_ids.shape[-1]:]
454
- final_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
455
- yield f"data: {json.dumps({'type': 'final', 'text': final_text, 'total_steps': 1})}\n\n"
456
-
457
- return Response(stream(), mimetype='text/event-stream', headers={'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'})
458
-
459
- @app.route('/')
460
- def index():
461
- return {
462
- "status": "healthy",
463
- "message": f"Multi-architecture API Router up and running. Target: {'Diffusion Framework' if IS_DIFFUSION else 'Causal Baseline Model'}", "model_loaded": MODEL_NAME}, 200
464
-
465
- if __name__ == '__main__':
466
- load_model()
467
- app.run(host='0.0.0.0', port=int(os.getenv('PORT', 7860)))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import sys
2
+ import types
3
+ import os
4
+ import math
5
+ import json
6
+ import torch
7
+ import torch.nn.functional as F
8
+ from flask import Flask, request, jsonify, Response
9
+ from transformers import AutoTokenizer, AutoModelForMaskedLM, AutoModelForCausalLM, TextIteratorStreamer
10
+ from threading import Thread
11
+
12
+ # 1. Environment Parsing & Architecture Strategy Mapping
13
+ MODEL_NAME = os.getenv("MODEL_NAME", "dllm-hub/Qwen3-0.6B-diffusion-bd3lm-v0.1")
14
+ IS_DIFFUSION = "diffusion" in MODEL_NAME.lower()
15
+
16
+ # Dynamic initialization layer targeting Diffusion Language Models
17
+ if IS_DIFFUSION:
18
+ try:
19
+ import dllm.utils
20
+ import dllm.pipelines
21
+ import dllm.data
22
+ import dllm.core
23
+ except ImportError:
24
+ pass
25
+ if 'dllm' not in sys.modules:
26
+ dllm_mock = types.ModuleType('dllm')
27
+ dllm_mock.core = sys.modules.get('dllm.core')
28
+ dllm_mock.data = sys.modules.get('dllm.data')
29
+ dllm_mock.pipelines = sys.modules.get('dllm.pipelines')
30
+ dllm_mock.utils = sys.modules.get('dllm.utils')
31
+ sys.modules['dllm'] = dllm_mock
32
+
33
+ app = Flask(__name__)
34
+ model = None
35
+ tokenizer = None
36
+ device = None
37
+
38
+ # ==========================================================
39
+ # SYSTEM WORKSPACE PIPELINES: CORE DIFFUSION SAMPLING LOOPS
40
+ # ==========================================================
41
+
42
+ def add_gumbel_noise(logits, temperature):
43
+ """Add Gumbel noise using float32 (faster than float64 on most GPUs)."""
44
+ if temperature == 0:
45
+ return logits
46
+ logits = logits.float()
47
+ noise = torch.rand_like(logits)
48
+ g = (-torch.log(noise)) ** temperature
49
+ return logits.exp() / g
50
+
51
+
52
+ def get_num_transfer_tokens(mask_index, steps):
53
+ mask_num = mask_index.sum(dim=1, keepdim=True)
54
+ base = mask_num // steps
55
+ rem = mask_num % steps
56
+ out = torch.zeros(mask_num.size(0), steps, device=mask_index.device, dtype=torch.long) + base
57
+ for i in range(mask_num.size(0)):
58
+ out[i, : rem[i]] += 1
59
+ return out
60
+
61
+
62
+ def build_staircase_attention_mask(x, block_size, pad_id):
63
+ B, T = x.shape
64
+ device = x.device
65
+ valid = x != pad_id
66
+ pos_raw = torch.cumsum(valid.long(), dim=-1)
67
+ position_ids = torch.where(valid, pos_raw - 1, torch.zeros_like(pos_raw)).long()
68
+ col = torch.arange(T, device=device)
69
+ block_ids = (col // block_size).view(1, T).expand(B, T)
70
+ block_ids = torch.where(valid, block_ids, torch.full_like(block_ids, -1))
71
+ q = block_ids.view(B, 1, T, 1)
72
+ k = block_ids.view(B, 1, 1, T)
73
+ attn = (k <= q) & (q >= 0) & (k >= 0)
74
+ return attn, position_ids
75
+
76
+
77
+ def clone_past_key_values(pkv):
78
+ """Shallow clone of prefix KV-cache to avoid expensive copy.deepcopy."""
79
+ if pkv is None:
80
+ return None
81
+ if isinstance(pkv, tuple):
82
+ return tuple(
83
+ (k.clone() if k is not None else None, v.clone() if v is not None else None)
84
+ for k, v in pkv
85
+ )
86
+ return pkv
87
+
88
+
89
+ def diffusion_step_block(logits, x_block, mask_block, num_transfer, temperature, remasking):
90
+ """Vectorized diffusion step — no per-sample Python loops."""
91
+ B, L, _ = logits.shape
92
+ if not mask_block.any():
93
+ return x_block
94
+ noisy = add_gumbel_noise(logits, temperature)
95
+ x0 = noisy.argmax(dim=-1)
96
+ if remasking == "low_confidence":
97
+ p = F.softmax(logits, dim=-1)
98
+ conf = p.gather(-1, x0.unsqueeze(-1)).squeeze(-1)
99
+ elif remasking == "random":
100
+ conf = torch.rand((B, L), device=logits.device)
101
+ else:
102
+ raise ValueError(remasking)
103
+ x0 = torch.where(mask_block, x0, x_block)
104
+ conf = conf.masked_fill(~mask_block, float("-inf"))
105
+ k_max = int(num_transfer.max().item())
106
+ if k_max > 0:
107
+ k = min(k_max, L)
108
+ topk_vals, topk_idx = torch.topk(conf, k=k, dim=-1)
109
+ commit = torch.zeros_like(x_block, dtype=torch.bool)
110
+ valid_mask = torch.arange(k, device=x_block.device).view(1, k) < num_transfer.view(B, 1)
111
+ commit.scatter_(1, topk_idx, valid_mask)
112
+ x_block = torch.where(commit, x0, x_block)
113
+ return x_block
114
+
115
+
116
+ @torch.inference_mode()
117
+ def generate(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size=32,
118
+ temperature=0.0, cfg_scale=0.0, remasking="low_confidence", capture_interval=0):
119
+ device = model.device
120
+ mask_id = tokenizer.mask_token_id
121
+ pad_id = tokenizer.pad_token_id
122
+ if pad_id is None:
123
+ pad_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else tokenizer.mask_token_id
124
+ if isinstance(prompt, torch.Tensor):
125
+ x = prompt.to(device).long()
126
+ else:
127
+ if isinstance(prompt[0], (list, tuple)):
128
+ max_len = max(len(p) for p in prompt)
129
+ x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
130
+ for i, p in enumerate(prompt):
131
+ x[i, : len(p)] = torch.tensor(p, device=device)
132
+ else:
133
+ x = torch.tensor(prompt, device=device).long()
134
+ if x.dim() == 1:
135
+ x = x.unsqueeze(0)
136
+ B = x.size(0)
137
+ finished = torch.zeros(B, dtype=torch.bool, device=device)
138
+ num_blocks = math.ceil(max_new_tokens / block_size)
139
+ steps_per_block = math.ceil(steps / num_blocks)
140
+ generated = 0
141
+ intermediates = []
142
+ total_step = 0
143
+ while generated < max_new_tokens:
144
+ if finished.all():
145
+ break
146
+ T_prefix = x.size(1)
147
+ offset = T_prefix % block_size
148
+ room = block_size if offset == 0 else block_size - offset
149
+ cur_len = min(room, max_new_tokens - generated)
150
+ if cur_len <= 0:
151
+ break
152
+ attn_pfx, pos_pfx = build_staircase_attention_mask(x, block_size, pad_id)
153
+ out = model(x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
154
+ cond_past = out.past_key_values
155
+ if cfg_scale > 0:
156
+ un_x = x.clone()
157
+ un_x[:] = mask_id
158
+ out_un = model(un_x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
159
+ uncond_past = out_un.past_key_values
160
+ else:
161
+ uncond_past = None
162
+ block = torch.full((B, cur_len), mask_id, device=device, dtype=torch.long)
163
+ block[finished] = pad_id
164
+ x = torch.cat([x, block], dim=1)
165
+ T_total = x.size(1)
166
+ block_mask = x[:, -cur_len:] == mask_id
167
+ num_transfer = get_num_transfer_tokens(block_mask, steps_per_block)
168
+ eff_steps = num_transfer.size(1)
169
+ full_attn, full_pos = build_staircase_attention_mask(x, block_size, pad_id)
170
+ attn_blk = full_attn[:, :, T_prefix:T_total, :]
171
+ pos_blk = full_pos[:, T_prefix:T_total]
172
+ # Clone prefix caches once per block instead of every step
173
+ cond_past_clone = clone_past_key_values(cond_past)
174
+ uncond_past_clone = clone_past_key_values(uncond_past) if uncond_past is not None else None
175
+ for t in range(eff_steps):
176
+ x_blk = x[:, T_prefix:T_total]
177
+ m_blk = x_blk == mask_id
178
+ cond_logits = model(
179
+ x_blk, attention_mask=attn_blk, position_ids=pos_blk,
180
+ past_key_values=cond_past_clone, use_cache=False
181
+ ).logits
182
+ logits = cond_logits
183
+ if cfg_scale > 0:
184
+ un_logits = model(
185
+ x_blk, attention_mask=attn_blk, position_ids=pos_blk,
186
+ past_key_values=uncond_past_clone, use_cache=False
187
+ ).logits
188
+ logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
189
+ x_blk_new = diffusion_step_block(
190
+ logits, x_blk, m_blk, num_transfer[:, t], temperature, remasking
191
+ )
192
+ x[:, T_prefix:T_total] = x_blk_new
193
+ if capture_interval > 0 and total_step % capture_interval == 0:
194
+ intermediates.append(x.clone())
195
+ total_step += 1
196
+ if tokenizer.eos_token_id is not None:
197
+ finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
198
+ generated += cur_len
199
+ if finished.all():
200
+ break
201
+ if capture_interval > 0:
202
+ return x, intermediates
203
+ return x
204
+
205
+
206
+ @torch.inference_mode()
207
+ def generate_stream(model, tokenizer, prompt, steps=128, max_new_tokens=128, block_size=32,
208
+ temperature=0.0, cfg_scale=0.0, remasking="low_confidence", capture_interval=10):
209
+ device = model.device
210
+ mask_id = tokenizer.mask_token_id
211
+ pad_id = tokenizer.pad_token_id
212
+ if pad_id is None:
213
+ pad_id = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else tokenizer.mask_token_id
214
+ if isinstance(prompt, torch.Tensor):
215
+ x = prompt.to(device).long()
216
+ else:
217
+ if isinstance(prompt[0], (list, tuple)):
218
+ max_len = max(len(p) for p in prompt)
219
+ x = torch.full((len(prompt), max_len), pad_id, device=device, dtype=torch.long)
220
+ for i, p in enumerate(prompt):
221
+ x[i, : len(p)] = torch.tensor(p, device=device)
222
+ else:
223
+ x = torch.tensor(prompt, device=device).long()
224
+ if x.dim() == 1:
225
+ x = x.unsqueeze(0)
226
+ B = x.size(0)
227
+ finished = torch.zeros(B, dtype=torch.bool, device=device)
228
+ num_blocks = math.ceil(max_new_tokens / block_size)
229
+ steps_per_block = math.ceil(steps / num_blocks)
230
+ generated = 0
231
+ total_step = 0
232
+ prompt_len = x.size(1)
233
+ while generated < max_new_tokens:
234
+ if finished.all():
235
+ break
236
+ T_prefix = x.size(1)
237
+ offset = T_prefix % block_size
238
+ room = block_size if offset == 0 else block_size - offset
239
+ cur_len = min(room, max_new_tokens - generated)
240
+ if cur_len <= 0:
241
+ break
242
+ attn_pfx, pos_pfx = build_staircase_attention_mask(x, block_size, pad_id)
243
+ out = model(x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
244
+ cond_past = out.past_key_values
245
+ if cfg_scale > 0:
246
+ un_x = x.clone()
247
+ un_x[:] = mask_id
248
+ out_un = model(un_x, attention_mask=attn_pfx, position_ids=pos_pfx, use_cache=True)
249
+ uncond_past = out_un.past_key_values
250
+ else:
251
+ uncond_past = None
252
+ block = torch.full((B, cur_len), mask_id, device=device, dtype=torch.long)
253
+ block[finished] = pad_id
254
+ x = torch.cat([x, block], dim=1)
255
+ T_total = x.size(1)
256
+ block_mask = x[:, -cur_len:] == mask_id
257
+ num_transfer = get_num_transfer_tokens(block_mask, steps_per_block)
258
+ eff_steps = num_transfer.size(1)
259
+ full_attn, full_pos = build_staircase_attention_mask(x, block_size, pad_id)
260
+ attn_blk = full_attn[:, :, T_prefix:T_total, :]
261
+ pos_blk = full_pos[:, T_prefix:T_total]
262
+ # Clone prefix caches once per block
263
+ cond_past_clone = clone_past_key_values(cond_past)
264
+ uncond_past_clone = clone_past_key_values(uncond_past) if uncond_past is not None else None
265
+ for t in range(eff_steps):
266
+ x_blk = x[:, T_prefix:T_total]
267
+ m_blk = x_blk == mask_id
268
+ cond_logits = model(
269
+ x_blk, attention_mask=attn_blk, position_ids=pos_blk,
270
+ past_key_values=cond_past_clone, use_cache=False
271
+ ).logits
272
+ logits = cond_logits
273
+ if cfg_scale > 0:
274
+ un_logits = model(
275
+ x_blk, attention_mask=attn_blk, position_ids=pos_blk,
276
+ past_key_values=uncond_past_clone, use_cache=False
277
+ ).logits
278
+ logits = un_logits + (cfg_scale + 1.0) * (cond_logits - un_logits)
279
+ x_blk_new = diffusion_step_block(
280
+ logits, x_blk, m_blk, num_transfer[:, t], temperature, remasking
281
+ )
282
+ x[:, T_prefix:T_total] = x_blk_new
283
+ if total_step % capture_interval == 0:
284
+ new_tokens = x[0, prompt_len:prompt_len + max_new_tokens].tolist()
285
+ text = tokenizer.decode(new_tokens, skip_special_tokens=True)
286
+ yield {
287
+ "type": "intermediate",
288
+ "step": total_step,
289
+ "text": text,
290
+ "total_steps": steps
291
+ }
292
+ total_step += 1
293
+ if tokenizer.eos_token_id is not None:
294
+ finished |= (x_blk_new == tokenizer.eos_token_id).any(dim=1)
295
+ if finished.all():
296
+ break
297
+ generated += cur_len
298
+ if finished.all():
299
+ break
300
+ new_tokens = x[0, prompt_len:prompt_len + max_new_tokens].tolist()
301
+ final_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
302
+ yield {
303
+ "type": "final",
304
+ "text": final_text,
305
+ "total_steps": total_step
306
+ }
307
+
308
+
309
+ # ==========================================================
310
+ # ARCHITECTURE ROUTING LAYERS & TRANSLATION ENGINE CODES
311
+ # ==========================================================
312
+
313
+ def load_model():
314
+ global model, tokenizer, device
315
+ device = "cuda" if torch.cuda.is_available() else "cpu"
316
+ print(f"Initializing {MODEL_NAME} on {device}... (Diffusion Strategy Flag = {IS_DIFFUSION})")
317
+ if IS_DIFFUSION:
318
+ model = AutoModelForMaskedLM.from_pretrained(
319
+ MODEL_NAME,
320
+ torch_dtype=torch.bfloat16,
321
+ trust_remote_code=True
322
+ ).to(device).eval()
323
+ else:
324
+ model = AutoModelForCausalLM.from_pretrained(
325
+ MODEL_NAME,
326
+ torch_dtype=torch.bfloat16,
327
+ trust_remote_code=False
328
+ ).to(device).eval()
329
+ # Compile model for faster inference (reduce-overhead is safest for generative loops)
330
+ try:
331
+ model = torch.compile(model, mode="reduce-overhead", fullgraph=False)
332
+ print("Model compiled with torch.compile.")
333
+ except Exception as e:
334
+ print(f"torch.compile skipped: {e}")
335
+ tokenizer = AutoTokenizer.from_pretrained(
336
+ MODEL_NAME,
337
+ trust_remote_code=IS_DIFFUSION
338
+ )
339
+ print("Model compilation completed and loaded into memory workspace.")
340
+
341
+
342
+ @app.route('/health', methods=['GET'])
343
+ def health():
344
+ return jsonify({"status": "healthy", "model_loaded": model is not None, "is_diffusion": IS_DIFFUSION})
345
+
346
+
347
+ @app.route('/generate', methods=['POST'])
348
+ def generate_text():
349
+ if model is None or tokenizer is None:
350
+ return jsonify({"error": "Model initialization missing"}), 503
351
+ data = request.get_json() or {}
352
+ if 'prompt' not in data:
353
+ return jsonify({"error": "Missing 'prompt' operational field"}), 400
354
+ prompt = data['prompt']
355
+ max_new_tokens = data.get('max_new_tokens', 256)
356
+ temperature = data.get('temperature', 0.0)
357
+ system_prompt = data.get('system_prompt', 'You are an expert assistant.')
358
+ messages = [
359
+ {"role": "system", "content": system_prompt},
360
+ {"role": "user", "content": prompt}
361
+ ]
362
+ encoded = tokenizer.apply_chat_template(
363
+ messages,
364
+ add_generation_prompt=True,
365
+ tokenize=True,
366
+ enable_thinking=False if IS_DIFFUSION else None
367
+ )
368
+ if IS_DIFFUSION:
369
+ input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
370
+ steps = data.get('steps', 256)
371
+ block_size = data.get('block_size', 32)
372
+ cfg_scale = data.get('cfg_scale', 0.0)
373
+ remasking = data.get('remasking', 'low_confidence')
374
+ output = generate(
375
+ model, tokenizer, input_ids,
376
+ steps=steps, max_new_tokens=max_new_tokens, block_size=block_size,
377
+ temperature=temperature, cfg_scale=cfg_scale, remasking=remasking,
378
+ )
379
+ prompt_len = len(encoded)
380
+ new_tokens = output[0, prompt_len:prompt_len + max_new_tokens].tolist()
381
+ generated_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
382
+ else:
383
+ input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
384
+ output_ids = model.generate(
385
+ input_ids,
386
+ max_new_tokens=max_new_tokens,
387
+ temperature=temperature,
388
+ do_sample=True if temperature > 0 else False,
389
+ pad_token_id=tokenizer.eos_token_id
390
+ )
391
+ generated_ids = output_ids[0, input_ids.shape[-1]:]
392
+ generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
393
+ return jsonify({"prompt": prompt, "generated_text": generated_text})
394
+
395
+
396
+ @app.route('/generate_stream', methods=['POST'])
397
+ def generate_text_stream():
398
+ if model is None or tokenizer is None:
399
+ return jsonify({"error": "Model workspace offline"}), 503
400
+ data = request.get_json() or {}
401
+ if not data or 'prompt' not in data:
402
+ return jsonify({"error": "Missing 'prompt' operational field"}), 400
403
+ prompt = data['prompt']
404
+ max_new_tokens = data.get('max_new_tokens', 256)
405
+ temperature = data.get('temperature', 0.0)
406
+ system_prompt = data.get('system_prompt', 'You are an expert assistant.')
407
+ messages = [
408
+ {"role": "system", "content": system_prompt},
409
+ {"role": "user", "content": prompt}
410
+ ]
411
+ encoded = tokenizer.apply_chat_template(
412
+ messages, add_generation_prompt=True, tokenize=True,
413
+ enable_thinking=False if IS_DIFFUSION else None
414
+ )
415
+ if IS_DIFFUSION:
416
+ input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
417
+ steps = data.get('steps', 256)
418
+ block_size = data.get('block_size', 32)
419
+ cfg_scale = data.get('cfg_scale', 0.0)
420
+ remasking = data.get('remasking', 'low_confidence')
421
+ capture_interval = data.get('capture_interval', 10)
422
+ output, intermediates = generate(
423
+ model, tokenizer, input_ids,
424
+ steps=steps, max_new_tokens=max_new_tokens, block_size=block_size,
425
+ temperature=temperature, cfg_scale=cfg_scale, remasking=remasking,
426
+ capture_interval=capture_interval,
427
+ )
428
+ prompt_len = len(encoded)
429
+ intermediate_states = []
430
+ for i, intermediate in enumerate(intermediates):
431
+ new_tokens = intermediate[0, prompt_len:prompt_len + max_new_tokens].tolist()
432
+ text = tokenizer.decode(new_tokens, skip_special_tokens=True)
433
+ intermediate_states.append({"step": i * capture_interval, "text": text})
434
+ new_tokens = output[0, prompt_len:prompt_len + max_new_tokens].tolist()
435
+ generated_text = tokenizer.decode(new_tokens, skip_special_tokens=True)
436
+ return jsonify({"prompt": prompt, "generated_text": generated_text, "intermediate_states": intermediate_states})
437
+ else:
438
+ input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
439
+ output_ids = model.generate(
440
+ input_ids, max_new_tokens=max_new_tokens, temperature=temperature,
441
+ do_sample=True if temperature > 0 else False,
442
+ pad_token_id=tokenizer.eos_token_id
443
+ )
444
+ generated_ids = output_ids[0, input_ids.shape[-1]:]
445
+ generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
446
+ return jsonify({"prompt": prompt, "generated_text": generated_text, "intermediate_states": []})
447
+
448
+
449
+ @app.route('/generate_sse', methods=['POST'])
450
+ def generate_text_sse():
451
+ if model is None or tokenizer is None:
452
+ return jsonify({"error": "Model workspace offline"}), 503
453
+ data = request.get_json() or {}
454
+ if not data or 'prompt' not in data:
455
+ return jsonify({"error": "Missing 'prompt' operational field"}), 400
456
+ prompt = data['prompt']
457
+ max_new_tokens = data.get('max_new_tokens', 256)
458
+ temperature = data.get('temperature', 0.0)
459
+ system_prompt = data.get('system_prompt', 'You are an expert real-time translation assistant.')
460
+ messages = [
461
+ {"role": "system", "content": system_prompt},
462
+ {"role": "user", "content": prompt}
463
+ ]
464
+ encoded = tokenizer.apply_chat_template(
465
+ messages, add_generation_prompt=True, tokenize=True,
466
+ enable_thinking=False if IS_DIFFUSION else None
467
+ )
468
+ input_ids = torch.tensor([encoded], dtype=torch.long, device=device)
469
+
470
+ def stream():
471
+ if IS_DIFFUSION:
472
+ steps = data.get('steps', 256)
473
+ block_size = data.get('block_size', 32)
474
+ cfg_scale = data.get('cfg_scale', 0.0)
475
+ remasking = data.get('remasking', 'low_confidence')
476
+ capture_interval = data.get('capture_interval', 10)
477
+ for state in generate_stream(
478
+ model, tokenizer, input_ids,
479
+ steps=steps, max_new_tokens=max_new_tokens, block_size=block_size,
480
+ temperature=temperature, cfg_scale=cfg_scale, remasking=remasking,
481
+ capture_interval=capture_interval
482
+ ):
483
+ yield f"data: {json.dumps(state)}\n\n"
484
+ else:
485
+ streamer = TextIteratorStreamer(
486
+ tokenizer, skip_prompt=True, skip_special_tokens=True
487
+ )
488
+ generation_kwargs = dict(
489
+ input_ids=input_ids,
490
+ streamer=streamer,
491
+ max_new_tokens=max_new_tokens,
492
+ temperature=temperature,
493
+ do_sample=True if temperature > 0 else False,
494
+ pad_token_id=tokenizer.eos_token_id,
495
+ )
496
+ def _generate():
497
+ with torch.inference_mode():
498
+ model.generate(**generation_kwargs)
499
+ thread = Thread(target=_generate)
500
+ thread.start()
501
+ for text in streamer:
502
+ yield f"data: {json.dumps({'type': 'token', 'text': text})}\n\n"
503
+ yield f"data: {json.dumps({'type': 'final', 'text': '', 'total_steps': 1})}\n\n"
504
+
505
+ return Response(
506
+ stream(), mimetype='text/event-stream',
507
+ headers={'Cache-Control': 'no-cache', 'X-Accel-Buffering': 'no'}
508
+ )
509
+
510
+
511
+ @app.route('/')
512
+ def index():
513
+ return {
514
+ "status": "healthy",
515
+ "message": f"Multi-architecture API Router up and running. Target: {'Diffusion Framework' if IS_DIFFUSION else 'Causal Baseline Model'}",
516
+ "model_loaded": MODEL_NAME
517
+ }, 200
518
+
519
+
520
+ if __name__ == '__main__':
521
+ load_model()
522
+ app.run(host='0.0.0.0', port=int(os.getenv('PORT', 7860)))