Andhs commited on
Commit
78590ea
·
verified ·
1 Parent(s): f37a7c9

Upload app.py

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