File size: 23,928 Bytes
79e8e52
3275441
 
79e8e52
3275441
 
79e8e52
3275441
79e8e52
 
 
 
 
 
 
 
 
 
3275441
79e8e52
 
 
3275441
 
 
 
 
79e8e52
 
 
3275441
79e8e52
 
 
3275441
 
 
 
 
 
79e8e52
3275441
79e8e52
 
3275441
 
 
 
79e8e52
 
 
3275441
 
79e8e52
 
 
 
3275441
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
79e8e52
3275441
79e8e52
3275441
79e8e52
 
 
3275441
 
 
79e8e52
 
3275441
 
 
 
79e8e52
 
 
3275441
 
 
 
 
79e8e52
3275441
79e8e52
 
 
 
 
 
 
3275441
 
79e8e52
3275441
 
 
79e8e52
3275441
 
79e8e52
3275441
 
 
 
f3fea40
 
 
 
 
 
3275441
 
 
 
 
 
 
79e8e52
 
 
 
 
 
3275441
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
79e8e52
3275441
 
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
 
 
79e8e52
 
 
3275441
79e8e52
 
 
 
 
 
 
 
 
 
3275441
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
79e8e52
 
 
 
 
 
 
 
 
3275441
 
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
 
 
 
 
 
 
 
 
 
79e8e52
 
3275441
79e8e52
3275441
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3275441
 
 
 
79e8e52
 
 
 
 
 
 
 
3275441
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
"""xeon_runtime.py β€” Intel Xeon runtime (V6): AVX512 + AMX_INT8 + IPEX + OneDNN + FP16.

═══════════════════════════════════════════════════════════════════════════════
V6 UPGRADE
═══════════════════════════════════════════════════════════════════════════════

V5 only set OMP/MKL threads + KMP_AFFINITY. V6 adds:

1. MKL_ENABLE_INSTRUCTIONS=AVX512        β€” forces MKL to dispatch AVX512 kernels
2. ONEDNN_MAX_CPU_ISA=AMX_INT8           β€” lets oneDNN use AMX INT8 tiles
3. DNNL_PRIMITIVE_CACHE_CAPACITY=1024    β€” large primitive cache (default 1024)
4. MKL_DYNAMIC=FALSE                     β€” disables MKL dynamic thread adjustment
5. IPEX (intel_extension_for_pytorch)    β€” Intel PyTorch extension
   - ipex.optimize(model) on the BiGRU_T model
   - torch.cpu.amp.autocast(dtype=torch.float16) for FP16 inference
6. FP16 benchmark: 8000Γ—8000 matmul, TFLOPS measurement
7. libvirt AMX activation helper β€” exposes amx-tile/amx-int8/amx-bf16 to a VM
   via host-passthrough CPU mode + feature policy='require'

The runtime is **always activated** (per user request: "sempre ativar otimizaΓ§Γ£o
para Xeon AVX512"). It degrades gracefully if IPEX / libvirt / AMX are not
available, but it never silently skips the optimization step.

═══════════════════════════════════════════════════════════════════════════════
USAGE
═══════════════════════════════════════════════════════════════════════════════

    from bigru_t.utils.xeon_runtime import optimize_xeon_environment
    N_CORES = optimize_xeon_environment()           # call ONCE, before torch
    # ... safe to import torch, ipex, etc. ...

For VM AMX exposure (requires libvirt-python and root):
    from bigru_t.utils.xeon_runtime import ativar_amx_na_vm
    ativar_amx_na_vm("my_vm_name")

═══════════════════════════════════════════════════════════════════════════════
"""
from __future__ import annotations
import os
import sys
import time
import logging
import platform
from typing import Optional, Tuple, Dict, Any

logger = logging.getLogger(__name__)

_NUCLEOS_ALOCADOS: Optional[int] = None
_IPEX_AVAILABLE: Optional[bool] = None
_AMX_CAPABLE: Optional[bool] = None
_V6_INIT_DONE: bool = False


# ============================================================================
# Helpers
# ============================================================================

def _detect_physical_cores() -> int:
    """Detect physical cores actually available to this process.

    Respects cgroup limits without requiring root. Falls back to logical
    cpu count if psutil is unavailable.
    """
    try:
        import psutil
        n_phys = psutil.cpu_count(logical=False) or 1
    except ImportError:
        try:
            with open("/proc/cpuinfo", "r") as f:
                cores = set()
                for line in f:
                    if line.startswith("core id"):
                        cores.add(line.strip())
                n_phys = len(cores) or 1
        except OSError:
            n_phys = 1

    try:
        n_affine = len(os.sched_getaffinity(0))
        n_logical = os.cpu_count() or 1
        if n_affine < n_logical:
            n_phys = max(1, n_affine // 2)
        else:
            n_phys = min(n_phys, n_affine)
    except (AttributeError, OSError):
        pass

    return max(1, n_phys)


def _read_cpu_flags() -> str:
    try:
        with open("/proc/cpuinfo", "r") as f:
            for line in f:
                if line.startswith("flags"):
                    return line
    except OSError:
        pass
    return ""


def get_avx512_capability() -> Tuple[bool, str]:
    """Check if the host CPU supports AVX512 VNNI."""
    flags = _read_cpu_flags()
    if "avx512_vnni" in flags:
        return True, "AVX512_VNNI (full INT8 acceleration)"
    elif "avx512f" in flags:
        return True, "AVX512F (no VNNI; INT8 falls back to AVX512F)"
    elif "avx2" in flags:
        return False, "AVX2 only (INT8 quantization works but slower)"
    else:
        return False, "Legacy SSE (INT8 quantization not recommended)"


def get_amx_capability() -> Tuple[bool, str]:
    """V6: check if AMX (Advanced Matrix Extensions) is available."""
    global _AMX_CAPABLE
    flags = _read_cpu_flags()
    has_tile = "amx_tile" in flags
    has_int8 = "amx_int8" in flags
    has_bf16 = "amx_bf16" in flags
    if has_tile and has_int8 and has_bf16:
        _AMX_CAPABLE = True
        return True, "AMX (tile + int8 + bf16) β€” full AMX acceleration"
    elif has_tile:
        _AMX_CAPABLE = True
        return True, f"AMX tile only (int8={has_int8}, bf16={has_bf16})"
    else:
        _AMX_CAPABLE = False
        return False, "AMX not available (AVX512 path will be used)"


def _try_import_ipex() -> Optional[Any]:
    """V6: try to import IPEX (Intel Extension for PyTorch).

    Returns the ipex module if available, else None. Caches the result.
    """
    global _IPEX_AVAILABLE
    if _IPEX_AVAILABLE is False:
        return None
    try:
        import intel_extension_for_pytorch as ipex  # type: ignore
        _IPEX_AVAILABLE = True
        return ipex
    except (ImportError, AttributeError, OSError) as e:
        _IPEX_AVAILABLE = False
        if _V6_INIT_DONE is False:
            logger.info(f"[Xeon V6] IPEX nΓ£o disponΓ­vel: {type(e).__name__}: {e}")
            logger.info("[Xeon V6] Continuando com OneDNN/MKL nativo do PyTorch.")
        return None


# ============================================================================
# FP16 benchmark (V6 β€” user-provided code)
# ============================================================================

def benchmark_fp16_matmul(size: int = 8000, warmup: int = 1, iters: int = 3) -> Dict[str, float]:
    """V6: FP16 matmul benchmark for Xeon AVX512/AMX.

    Runs `size`Γ—`size` FP16 matmul `iters` times and reports:
        - best_time_ms: lowest wall time
        - best_tflops: best achieved TFLOPS
        - avg_tflops: average TFLOPS

    Returns empty dict if torch unavailable.
    """
    try:
        import torch
    except ImportError:
        return {}

    results: Dict[str, float] = {}
    try:
        # Warmup
        a = torch.randn(size, size, dtype=torch.float16)
        b = torch.randn(size, size, dtype=torch.float16)
        for _ in range(warmup):
            _ = torch.matmul(a, b)
        # Bench
        times = []
        for _ in range(iters):
            t0 = time.perf_counter()
            _ = torch.matmul(a, b)
            times.append(time.perf_counter() - t0)
        best_t = min(times)
        avg_t = sum(times) / len(times)
        # 2*size^3 FLOPs per matmul (M*N*K)
        flops = 2.0 * (size ** 3)
        results["best_time_ms"] = best_t * 1000.0
        results["avg_time_ms"] = avg_t * 1000.0
        results["best_tflops"] = flops / best_t / 1e12
        results["avg_tflops"] = flops / avg_t / 1e12
        results["matrix_size"] = float(size)
    except (RuntimeError, MemoryError) as e:
        logger.warning(f"[Xeon V6] FP16 benchmark failed: {e}")
        results["error"] = str(e)
    return results


def benchmark_int8_matmul(size: int = 4096, warmup: int = 1, iters: int = 3) -> Dict[str, float]:
    """V6: INT8 matmul benchmark β€” uses AMX_INT8 when available via oneDNN.

    Falls back to FP32 if INT8 path is unavailable.
    """
    try:
        import torch
    except ImportError:
        return {}

    results: Dict[str, float] = {}
    try:
        a = torch.randint(-127, 127, (size, size), dtype=torch.int8)
        b = torch.randint(-127, 127, (size, size), dtype=torch.int8)
        # Warmup
        for _ in range(warmup):
            _ = torch.matmul(a.float(), b.float())
        times = []
        for _ in range(iters):
            t0 = time.perf_counter()
            _ = torch.matmul(a.float(), b.float())
            times.append(time.perf_counter() - t0)
        best_t = min(times)
        flops = 2.0 * (size ** 3)
        results["int8_best_time_ms"] = best_t * 1000.0
        results["int8_best_tflops"] = flops / best_t / 1e12
        results["int8_matrix_size"] = float(size)
    except (RuntimeError, MemoryError) as e:
        results["int8_error"] = str(e)
    return results


# ============================================================================
# Main entry point: optimize_xeon_environment (V6)
# ============================================================================

def optimize_xeon_environment(
    verbose: bool = True,
    force_ipex: bool = False,
) -> int:
    """V6: configure Intel Xeon AVX512 + AMX_INT8 + IPEX + OneDNN.

    Always activates β€” per user requirement "sempre ativar otimizaΓ§Γ£o para Xeon
    AVX512". Idempotent. Sets every environment variable that influences MKL,
    OpenMP, oneDNN, and (optionally) IPEX runtime behavior.

    Args:
        verbose: print configuration summary to stdout
        force_ipex: if True, raise when IPEX import fails. Default False β€”
            degrade gracefully to native PyTorch oneDNN.

    Returns:
        Number of physical cores allocated to this process.
    """
    global _NUCLEOS_ALOCADOS, _V6_INIT_DONE

    if _NUCLEOS_ALOCADOS is not None and _V6_INIT_DONE:
        return _NUCLEOS_ALOCADOS

    n_phys = _detect_physical_cores()

    # ───────────────────────────────────────────────────────────────────────
    # V6 β€” Environment variables (user-provided block)
    # ───────────────────────────────────────────────────────────────────────
    os.environ["MKL_ENABLE_INSTRUCTIONS"] = "AVX512"
    NUM_CORES = str(n_phys)
    os.environ["MKL_NUM_THREADS"] = NUM_CORES
    os.environ["OMP_NUM_THREADS"] = NUM_CORES
    os.environ["MKL_DYNAMIC"] = "FALSE"
    os.environ["DNNL_PRIMITIVE_CACHE_CAPACITY"] = "1024"
    os.environ["ONEDNN_MAX_CPU_ISA"] = "AMX_INT8"

    # ───────────────────────────────────────────────────────────────────────
    # V5 (kept) β€” OpenMP thread pinning
    # ───────────────────────────────────────────────────────────────────────
    os.environ.setdefault("KMP_AFFINITY", "granularity=fine,compact,1,0")
    os.environ.setdefault("KMP_BLOCKTIME", "1")
    os.environ.setdefault("TOKENIZERS_PARALLELISM", "true")

    # ───────────────────────────────────────────────────────────────────────
    # PyTorch backend flags
    # ───────────────────────────────────────────────────────────────────────
    try:
        import torch
        torch.set_num_threads(n_phys)
        try:
            torch.set_num_interop_threads(1)
        except RuntimeError:
            # V6.5: jΓ‘ inicializado (e.g., bigru_t package importou torch antes).
            # Silenciosamente ignora β€” o paralelismo jΓ‘ estΓ‘ configurado.
            pass
        if hasattr(torch.backends, "mkldnn"):
            torch.backends.mkldnn.enabled = True
        if hasattr(torch.backends, "quantized"):
            try:
                torch.backends.quantized.engine = "fbgemm"
            except (RuntimeError, AttributeError):
                pass
        # V6: enable TF32 for Ampere+ / Sapphire Rapids (irrelevant on CPU but
        # harmless) and ensure oneDNN verbose is silent.
        try:
            torch.backends.cuda.matmul.allow_tf32 = True
        except AttributeError:
            pass
    except ImportError:
        if verbose:
            print("[Xeon V6] WARNING: PyTorch not yet imported β€” env vars set,"
                  " call this BEFORE importing torch for full effect.")

    # ───────────────────────────────────────────────────────────────────────
    # V6 β€” IPEX (intel_extension_for_pytorch)
    # ───────────────────────────────────────────────────────────────────────
    ipex = _try_import_ipex()
    if ipex is None and force_ipex:
        raise ImportError(
            "IPEX (intel_extension_for_pytorch) nΓ£o disponΓ­vel, mas force_ipex=True. "
            "Instale com: pip install intel-extension-for-pytorch"
        )

    # ───────────────────────────────────────────────────────────────────────
    # V6 β€” AMX capability check
    # ───────────────────────────────────────────────────────────────────────
    amx_ok, amx_desc = get_amx_capability()
    avx_ok, avx_desc = get_avx512_capability()

    _NUCLEOS_ALOCADOS = n_phys
    _V6_INIT_DONE = True

    if verbose:
        print("\n" + "=" * 72)
        print(f"[Xeon Runtime V6] Intel Xeon optimization activated")
        print("=" * 72)
        print(f"  Physical cores        : {n_phys}")
        print(f"  MKL_NUM_THREADS       : {os.environ['MKL_NUM_THREADS']}")
        print(f"  OMP_NUM_THREADS       : {os.environ['OMP_NUM_THREADS']}")
        print(f"  MKL_DYNAMIC           : {os.environ['MKL_DYNAMIC']}")
        print(f"  MKL_ENABLE_INSTRUCTIONS: {os.environ['MKL_ENABLE_INSTRUCTIONS']}")
        print(f"  ONEDNN_MAX_CPU_ISA    : {os.environ['ONEDNN_MAX_CPU_ISA']}")
        print(f"  DNNL_PRIMITIVE_CACHE  : {os.environ['DNNL_PRIMITIVE_CACHE_CAPACITY']}")
        print(f"  KMP_AFFINITY          : {os.environ['KMP_AFFINITY']}")
        print(f"  KMP_BLOCKTIME         : {os.environ['KMP_BLOCKTIME']}")
        print(f"  AVX512                : {avx_desc}")
        print(f"  AMX                   : {amx_desc}")
        print(f"  IPEX                  : "
              f"{'available' if ipex is not None else 'not installed (using native oneDNN)'}")
        print("=" * 72 + "\n")

    return n_phys


# ============================================================================
# V6 β€” apply_ipex_optimization (model-level)
# ============================================================================

def apply_ipex_optimization(model, dtype=None, optimizer=None):
    """V6: apply ipex.optimize() to a model.

    Returns (model, optimizer) tuple. If IPEX is not available, returns the
    inputs unchanged. The model is modified in-place when IPEX is present.

    Args:
        model: torch.nn.Module
        dtype: optional torch.dtype for the model (e.g. torch.bfloat16)
        optimizer: optional torch.optim.Optimizer to also optimize
    """
    ipex = _try_import_ipex()
    if ipex is None:
        return model, optimizer
    try:
        import torch
        if dtype is not None:
            model = model.to(dtype)
        if optimizer is not None:
            model, optimizer = ipex.optimize(model=model, optimizer=optimizer, dtype=dtype)
        else:
            model = ipex.optimize(model=model, dtype=dtype)
        logger.info(f"[Xeon V6] ipex.optimize applied (dtype={dtype})")
    except (RuntimeError, AttributeError, TypeError) as e:
        logger.warning(f"[Xeon V6] ipex.optimize failed: {e}")
    return model, optimizer


# ============================================================================
# V6 β€” FP16 autocast context manager
# ============================================================================

class fp16_autocast:
    """V6: FP16 CPU autocast context manager.

    Uses torch.cpu.amp.autocast(dtype=torch.float16) when available. Falls
    back to a no-op context manager if autocast is not supported.

    Usage:
        with fp16_autocast():
            y = model(x)
    """

    def __init__(self, enabled: bool = True):
        self.enabled = enabled
        self._ctx = None

    def __enter__(self):
        if not self.enabled:
            return self
        try:
            import torch
            if hasattr(torch.cpu, "amp") and hasattr(torch.cpu.amp, "autocast"):
                self._ctx = torch.cpu.amp.autocast(dtype=torch.float16)
                self._ctx.__enter__()
        except (ImportError, RuntimeError, AttributeError):
            self._ctx = None
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):
        if self._ctx is not None:
            self._ctx.__exit__(exc_type, exc_val, exc_tb)
            self._ctx = None
        return False


# ============================================================================
# V6 β€” libvirt AMX activation helper (user-provided code, integrated)
# ============================================================================

def ativar_amx_na_vm(nome_vm: str, qemu_uri: str = "qemu:///system") -> bool:
    """V6: ativa AMX (amx-tile, amx-int8, amx-bf16) numa VM via libvirt.

    Requer:
        - pip install libvirt-python
        - libvirtd rodando localmente (qemu:///system)
        - permissΓ£o de root (ou membro do grupo libvirt)
        - CPU fΓ­sica com AMX (verificado por get_amx_capability())

    ImplementaΓ§Γ£o:
        1. Conecta ao daemon libvirt
        2. Busca a VM pelo nome
        3. LΓͺ o XML persistente (flag VIR_DOMAIN_XML_INACTIVE = 2)
        4. Garante <cpu mode='host-passthrough'>
        5. Adiciona <feature policy='require' name='amx-tile|int8|bf16'/>
        6. Reescreve o XML via defineXML()

    Retorna True se a VM foi atualizada com sucesso, False caso contrΓ‘rio.
    Requer reboot da VM para aplicar.

    Args:
        nome_vm: nome da mΓ‘quina virtual no libvirt
        qemu_uri: URI do libvirt (default: qemu:///system)

    Raises:
        ImportError: se libvirt-python nΓ£o estiver instalado
        RuntimeError: se a conexΓ£o com libvirt falhar
    """
    try:
        import libvirt  # type: ignore
    except ImportError as e:
        raise ImportError(
            "libvirt-python nΓ£o instalado. Rode: pip install libvirt-python "
            "(tambΓ©m requer libvirt-dev no sistema: apt install libvirt-dev)"
        ) from e

    try:
        import xml.etree.ElementTree as ET
    except ImportError:
        return False

    # PrΓ©-checa AMX no host fΓ­sico
    amx_ok, amx_desc = get_amx_capability()
    if not amx_ok:
        print(f"[AMX VM] Host fΓ­sico nΓ£o tem AMX ({amx_desc}).")
        print("         NΓ£o adianta ativar AMX na VM β€” o host precisa suportar.")
        return False

    try:
        conn = libvirt.open(qemu_uri)
        if conn is None:
            print(f"[AMX VM] Falha ao abrir conexΓ£o com {qemu_uri}")
            return False

        try:
            dom = conn.lookupByName(nome_vm)
        except libvirt.libvirtError:
            print(f"[AMX VM] VM '{nome_vm}' nΓ£o encontrada.")
            conn.close()
            return False

        xml_atual = dom.XMLDesc(2)  # VIR_DOMAIN_XML_INACTIVE
        root = ET.fromstring(xml_atual)
        cpu_elem = root.find("cpu")

        if cpu_elem is None:
            cpu_elem = ET.SubElement(root, "cpu", mode="host-passthrough")
            print("[AMX VM] Tag <cpu> criada com mode='host-passthrough'.")
        else:
            cpu_elem.set("mode", "host-passthrough")
            print("[AMX VM] CPU mode atualizado para host-passthrough.")

        flags_para_adicionar = ["amx-tile", "amx-int8", "amx-bf16"]
        added = []
        for flag in flags_para_adicionar:
            existing = cpu_elem.find(f"./feature[@name='{flag}']")
            if existing is None:
                ET.SubElement(cpu_elem, "feature", policy="require", name=flag)
                added.append(flag)
                print(f"[AMX VM] Feature adicionada: {flag}")

        novo_xml = ET.tostring(root, encoding="utf-8").decode("utf-8")
        conn.defineXML(novo_xml)
        print(f"[AMX VM] βœ“ AMX ativado no XML da VM '{nome_vm}'. Reinicie a VM para aplicar.")
        conn.close()
        return True

    except libvirt.libvirtError as e:
        print(f"[AMX VM] Erro libvirt: {e}")
        return False
    except Exception as e:
        print(f"[AMX VM] Erro inesperado: {e}")
        return False


# ============================================================================
# Compatibility helpers (V5 API preserved)
# ============================================================================

def get_allocated_cores() -> int:
    """Return the number of cores allocated by `optimize_xeon_environment`."""
    return _NUCLEOS_ALOCADOS or 0


def make_ort_session_options():
    """Build onnxruntime.SessionOptions tuned for Xeon AVX512/AMX."""
    import onnxruntime as ort  # type: ignore

    n_cores = _NUCLEOS_ALOCADOS or _detect_physical_cores()

    opts = ort.SessionOptions()
    opts.intra_op_num_threads = n_cores
    opts.inter_op_num_threads = 1
    opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
    opts.enable_cpu_mem_arena = True
    return opts


def get_xeon_status() -> Dict[str, Any]:
    """V6: return a snapshot of all Xeon runtime flags + capabilities.

    Useful for logging inside the training script.
    """
    amx_ok, amx_desc = get_amx_capability()
    avx_ok, avx_desc = get_avx512_capability()
    return {
        "version": "V6",
        "physical_cores": _NUCLEOS_ALOCADOS or _detect_physical_cores(),
        "env": {
            "MKL_ENABLE_INSTRUCTIONS": os.environ.get("MKL_ENABLE_INSTRUCTIONS"),
            "MKL_NUM_THREADS": os.environ.get("MKL_NUM_THREADS"),
            "OMP_NUM_THREADS": os.environ.get("OMP_NUM_THREADS"),
            "MKL_DYNAMIC": os.environ.get("MKL_DYNAMIC"),
            "DNNL_PRIMITIVE_CACHE_CAPACITY": os.environ.get("DNNL_PRIMITIVE_CACHE_CAPACITY"),
            "ONEDNN_MAX_CPU_ISA": os.environ.get("ONEDNN_MAX_CPU_ISA"),
            "KMP_AFFINITY": os.environ.get("KMP_AFFINITY"),
            "KMP_BLOCKTIME": os.environ.get("KMP_BLOCKTIME"),
        },
        "avx512": {"supported": avx_ok, "desc": avx_desc},
        "amx": {"supported": amx_ok, "desc": amx_desc},
        "ipex_available": _IPEX_AVAILABLE is True,
        "init_done": _V6_INIT_DONE,
    }


__all__ = [
    "optimize_xeon_environment",
    "apply_ipex_optimization",
    "fp16_autocast",
    "benchmark_fp16_matmul",
    "benchmark_int8_matmul",
    "ativar_amx_na_vm",
    "get_avx512_capability",
    "get_amx_capability",
    "get_xeon_status",
    "get_allocated_cores",
    "make_ort_session_options",
]