print(“\n” + “=”*70 + “\n4. Lote empaquetado de longitud variable, sin desperdicio de relleno\n” + “=”*70) seqlens = [37, 120, 8, 200]
total = suma(seqlens) H, K = 8, 64 q = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) k = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) v = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) intente: sesgo = ab.BlockDiagonalMask.from_seqlens(seqlens) out_packed = xops.memory_ficient_attention(q, k, v, attn_bias=bias) s0 = seqlens[0]
ref0 = atención_vainilla(q[:, :s0]k[:, :s0]v[:, :s0]).half() print(“forma empaquetada:”, tupla(out_packed.shape), “(todos”, total, “tokens, sin pad)”) print(“segmento-0 max diff: {:.2e}”.format((out_packed)[:, :s0] – ref0).abs().max().item())) cbias = ab.BlockDiagonalCausalMask.from_seqlens(seqlens) _ = xops.memory_ficient_attention(q, k, v, attn_bias=cbias) print(“-> también hizo un pase CAUSAL empaquetado. Así es como los motores estilo vLLM”) print(” solicitudes por lotes de diferentes longitudes con sobrecarga de relleno cero.”) divisiones = sesgo.split(out_packed) print(“segmentos recuperados:”, [tuple(t.shape) for t in splits]) excepto excepción como e: print(“Ruta de BlockDiagonalMask omitida en esta versión/backend:”, repr(e)) print(“\n” + “=”*70 + “\n5. Atención de consultas agrupadas (diseño BMGHK 5-D)\n” + “=”*70) B, M, K = 2, 256, 64 n_q_heads, n_kv_heads = 8, 2 G, Hq = n_kv_heads, n_q_heads // n_kv_heads prueba: qg = torch.randn(B, M, G, Hq, K, dispositivo=dispositivo, dtype=torch.float16) kg = torch.randn(B, M, G, 1, K, dispositivo=dispositivo, dtype=torch.float16) vg = torch.randn(B, M, G, 1, K, dispositivo=dispositivo, dtype=torch.float16) out_gqa = xops.memory_ficient_attention(qg, kg, vg) print(“Forma de salida GQA:”, tupla(out_gqa.shape), “= [B, M, G, Hq, K]”) print(f”-> {n_q_heads} cabezales de consulta, solo {n_kv_heads} cabezales KV: caché KV más pequeño,”) print(” que es exactamente lo que los modelos clase Llama/Mistral usan en la inferencia.”) excepto excepción como e: print(“Ruta GQA 5-D omitida en esta versión/backend:”, repr(e))
total = suma(seqlens) H, K = 8, 64 q = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) k = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) v = torch.randn(1, total, H, K, dispositivo=dispositivo, dtype=torch.float16) intente: sesgo = ab.BlockDiagonalMask.from_seqlens(seqlens) out_packed = xops.memory_ficient_attention(q, k, v, attn_bias=bias) s0 = seqlens[0]
ref0 = atención_vainilla(q[:, :s0]k[:, :s0]v[:, :s0]).half() print(“forma empaquetada:”, tupla(out_packed.shape), “(todos”, total, “tokens, sin pad)”) print(“segmento-0 max diff: {:.2e}”.format((out_packed)[:, :s0] – ref0).abs().max().item())) cbias = ab.BlockDiagonalCausalMask.from_seqlens(seqlens) _ = xops.memory_ficient_attention(q, k, v, attn_bias=cbias) print(“-> también hizo un pase CAUSAL empaquetado. Así es como los motores estilo vLLM”) print(” solicitudes por lotes de diferentes longitudes con sobrecarga de relleno cero.”) divisiones = sesgo.split(out_packed) print(“segmentos recuperados:”, [tuple(t.shape) for t in splits]) excepto excepción como e: print(“Ruta de BlockDiagonalMask omitida en esta versión/backend:”, repr(e)) print(“\n” + “=”*70 + “\n5. Atención de consultas agrupadas (diseño BMGHK 5-D)\n” + “=”*70) B, M, K = 2, 256, 64 n_q_heads, n_kv_heads = 8, 2 G, Hq = n_kv_heads, n_q_heads // n_kv_heads prueba: qg = torch.randn(B, M, G, Hq, K, dispositivo=dispositivo, dtype=torch.float16) kg = torch.randn(B, M, G, 1, K, dispositivo=dispositivo, dtype=torch.float16) vg = torch.randn(B, M, G, 1, K, dispositivo=dispositivo, dtype=torch.float16) out_gqa = xops.memory_ficient_attention(qg, kg, vg) print(“Forma de salida GQA:”, tupla(out_gqa.shape), “= [B, M, G, Hq, K]”) print(f”-> {n_q_heads} cabezales de consulta, solo {n_kv_heads} cabezales KV: caché KV más pequeño,”) print(” que es exactamente lo que los modelos clase Llama/Mistral usan en la inferencia.”) excepto excepción como e: print(“Ruta GQA 5-D omitida en esta versión/backend:”, repr(e))