fix: fix MLA for ShardedModelLoader/RemoteModelLoader (#6287)
Signed-off-by: wangyu <wangyu.steph@bytedance.com>
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user