fix: fix MLA for ShardedModelLoader/RemoteModelLoader (#6287)

Signed-off-by: wangyu <wangyu.steph@bytedance.com>
This commit is contained in:
wangyu
2025-08-28 16:10:09 -07:00
committed by GitHub
parent a38c149758
commit 9f81d741a2
8 changed files with 37 additions and 35 deletions
+1 -1
View File
@@ -20,7 +20,7 @@ class ConnectorType(str, enum.Enum):
KV = "KV"
def create_remote_connector(url, device="cpu") -> BaseConnector:
def create_remote_connector(url, **kwargs) -> BaseConnector:
connector_type = parse_connector_type(url)
if connector_type == "redis":
return RedisConnector(url)
@@ -20,9 +20,8 @@ class BaseConnector(ABC):
<connector_type://<host>:<port>/<model_name>/files/<filename>
"""
def __init__(self, url: str, device: torch.device = "cpu"):
def __init__(self, url: str):
self.url = url
self.device = device
self.closed = False
self.local_dir = tempfile.mkdtemp()
for sig in (signal.SIGINT, signal.SIGTERM):
+2 -2
View File
@@ -15,10 +15,10 @@ logger = logging.getLogger(__name__)
class RedisConnector(BaseKVConnector):
def __init__(self, url: str, device: torch.device = "cpu"):
def __init__(self, url: str):
import redis
super().__init__(url, device)
super().__init__(url)
parsed_url = urlparse(url)
self.connection = redis.Redis(host=parsed_url.hostname, port=parsed_url.port)
self.model_name = parsed_url.path.lstrip("/")
@@ -15,7 +15,7 @@ def create_serde(serde_type: str) -> Tuple[Serializer, Deserializer]:
if serde_type == "safe":
s = SafeSerializer()
d = SafeDeserializer(torch.uint8)
d = SafeDeserializer()
else:
raise ValueError(f"Unknown serde type: {serde_type}")
@@ -19,11 +19,12 @@ class SafeSerializer(Serializer):
class SafeDeserializer(Deserializer):
def __init__(self, dtype):
super().__init__(dtype)
def __init__(self):
# TODO: dtype options
super().__init__(torch.float32)
def from_bytes_normal(self, b: Union[bytearray, bytes]) -> torch.Tensor:
return load(bytes(b))["tensor_bytes"].to(dtype=self.dtype)
return load(bytes(b))["tensor_bytes"]
def from_bytes(self, b: Union[bytearray, bytes]) -> torch.Tensor:
return self.from_bytes_normal(b)