Import and enable FlagGems#
FlagGems supports two common usage patterns: patching PyTorch ATen operators (recommended) and calling FlagGems operators explicitly.
Enable FlagGems and patch ATen operators After
flag_gems.enable(), supportedtorch.* / torch.nn.functional.*calls will be dispatched to FlagGems implementations automatically.Global enablement To apply FlagGems optimizations across your entire script or interactive session:
import torch import flag_gems flag_gems.enable() x = torch.randn(4096, 4096, device=flag_gems.device, dtype=torch.float16) y = torch.mm(x, x)
Once enabled, all supported operators in your code will automatically be replaced with the optimized
FlagGemsimplementations, no further changes needed.Selective enablement To enable only specific operators and skip the rest:
import flag_gems # Enable only selected operators flag_gems.only_enable(include=["rms_norm", "softmax"])
Scoped enablement For finer control, you can enable
FlagGemsonly within a specific code block or scope using a context manager:import torch import flag_gems with flag_gems.use_gems(): x = torch.randn(4096, 4096, device=flag_gems.device, dtype=torch.float16) y = torch.mm(x, x)
Enabling within a specific scope is helpful for the following cases:
Benchmark performance differences
Compare correctness between implementations
Apply acceleration selectively in complex workflows
Within the enabled scope, you can also selectively enable the operators in the context manager:
# Enable only specific operators in the scope with flag_gems.use_gems(include=["sum", "add"]): # Only sum and add will be accelerated ... # Or exclude specific operators with flag_gems.use_gems(exclude=["mul", "div"]): # All except mul and div will be accelerated ...
Note
The
includeparameter has higher priority thanexclude. If both are provided,excludeis ignored.
The
flag_gems.enable(...)andflag_gems.only_enable(...)functions support several optional parameters. For more information, see Use optional parameters for FlagGems enablement function.Explicitly call FlagGems ops You can also bypass PyTorch dispatch and call operators from
flag_gems.opsdirectly without usingenable():import torch from flag_gems import ops import flag_gems a = torch.randn(1024, 1024, device=flag_gems.device, dtype=torch.float16) b = torch.randn(1024, 1024, device=flag_gems.device, dtype=torch.float16) c = ops.mm(a, b)