| from __future__ import annotations |
| from contextlib import contextmanager |
| import gc |
| from types import SimpleNamespace |
| import torch |
| from diffusers.models.modeling_outputs import Transformer2DModelOutput |
| from trt_torch import TorchTensorRTEngine |
|
|
| class Flux2TRTTransformer(torch.nn.Module): |
| """Static 512 FLUX.2 Klein transformer backed by an L4 TensorRT plan.""" |
| def __init__(self, engine_path, config): |
| super().__init__() |
| self.engine_path=str(engine_path) |
| self.engine=None |
| |
| |
| self._offload_anchor=torch.nn.Parameter(torch.empty(0),requires_grad=False) |
| self.config=SimpleNamespace(**dict(config)) if isinstance(config,dict) else config |
| |
| |
| self._dtype=torch.bfloat16 |
| self._device=torch.device('cuda') |
| @property |
| def dtype(self): return self._dtype |
| @property |
| def device(self): return self._offload_anchor.device |
| @contextmanager |
| def cache_context(self, *args, **kwargs): |
| yield |
| def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, img_ids=None, txt_ids=None, |
| guidance=None, joint_attention_kwargs=None, return_dict=True, **kwargs): |
| if guidance is not None: |
| raise RuntimeError('Static Companion Forge FLUX.2 TRT engine is distilled/no-guidance only') |
| feeds={ |
| 'hidden_states':hidden_states.to(device='cuda',dtype=torch.float16), |
| 'encoder_hidden_states':encoder_hidden_states.to(device='cuda',dtype=torch.float16), |
| 'timestep':timestep.to(device='cuda',dtype=torch.float16), |
| 'img_ids':img_ids.to(device='cuda',dtype=torch.int64), |
| 'txt_ids':txt_ids.to(device='cuda',dtype=torch.int64), |
| } |
| if self.engine is None: |
| self.engine=TorchTensorRTEngine(self.engine_path) |
| y=self.engine.run(feeds)['sample'] |
| y=y.to(dtype=hidden_states.dtype) |
| return Transformer2DModelOutput(sample=y) if return_dict else (y,) |
| def to(self,*args,**kwargs): |
| |
| |
| device=kwargs.get('device', None) |
| if device is None and args: |
| first=args[0] |
| if isinstance(first,(str,torch.device,int)): device=first |
| if device is not None: |
| dev=torch.device(f'cuda:{device}' if isinstance(device,int) else device) |
| self._offload_anchor.data=self._offload_anchor.data.to(dev) |
| self._device=dev |
| if dev.type=='cpu' and self.engine is not None: |
| self.engine=None; gc.collect() |
| if torch.cuda.is_available(): torch.cuda.empty_cache() |
| return self |
| def cuda(self,device=None): return self.to(torch.device('cuda' if device is None else f'cuda:{device}')) |
| def cpu(self): return self.to(torch.device('cpu')) |
| def eval(self): super().eval(); return self |
|
|