[diffusion] fix: fix RowParallel LoRA merged forwarding (#24410)
This commit is contained in:
@@ -513,6 +513,9 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
|
||||
super().__init__(base_layer, lora_rank, lora_alpha)
|
||||
|
||||
def forward(self, input_: torch.Tensor):
|
||||
if self.merged or self.disable_lora:
|
||||
return self.base_layer(input_)
|
||||
|
||||
lora_A = self.lora_A
|
||||
lora_B = self.lora_B
|
||||
if isinstance(self.lora_B, DTensor):
|
||||
|
||||
@@ -33,7 +33,7 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
SGL_TEST_FILES_CI_DATA_REVISION = "437539b330592b1421d239b22d7d76b4a6d08fda"
|
||||
SGL_TEST_FILES_CI_DATA_REVISION = "4d9eff3b05b0ffe1d3529e8bb148b63af11a4b92"
|
||||
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
|
||||
"https://raw.githubusercontent.com/"
|
||||
f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/"
|
||||
|
||||
Reference in New Issue
Block a user