@@ -491,88 +491,88 @@ IoTDB> show models
491491
4924921 . AINode 目前使用 v4 .56 .2 版本的 transformers,构建模型时需** 避免继承低版本(< 4 .50 )接口** ;
4934932 . 模型需继承一类 AINode 的推理任务流水线(当前支持预测流水线):
494- 1 . iotdb- core/ ainode/ iotdb/ ainode/ core/ inference/ pipeline/ basic\_pipeline .py
495-
496- ` ` ` Python
497- class BasicPipeline(ABC):
498- def __init__(self, model_id, **model_kwargs):
499- self.model_info = model_info
500- self.device = model_kwargs.get("device", "cpu")
501- self.model = load_model(model_info, device_map=self.device, **model_kwargs)
494+ * iotdb- core/ ainode/ iotdb/ ainode/ core/ inference/ pipeline/ basic\_pipeline .py
495+
496+ ` ` ` Python
497+ class BasicPipeline(ABC):
498+ def __init__(self, model_id, **model_kwargs):
499+ self.model_info = model_info
500+ self.device = model_kwargs.get("device", "cpu")
501+ self.model = load_model(model_info, device_map=self.device, **model_kwargs)
502502
503- @abstractmethod
504- def preprocess(self, inputs, **infer_kwargs):
505- """
506- 在推理任务开始前对输入数据进行前处理,包括形状验证和数值转换。
507- """
508- pass
503+ @abstractmethod
504+ def preprocess(self, inputs, **infer_kwargs):
505+ """
506+ 在推理任务开始前对输入数据进行前处理,包括形状验证和数值转换。
507+ """
508+ pass
509509
510- @abstractmethod
511- def postprocess(self, output, **infer_kwargs):
512- """
513- 在推理任务结束后对输出结果进行后处理。
514- """
515- pass
510+ @abstractmethod
511+ def postprocess(self, output, **infer_kwargs):
512+ """
513+ 在推理任务结束后对输出结果进行后处理。
514+ """
515+ pass
516516
517517
518- class ForecastPipeline(BasicPipeline):
519- def __init__(self, model_info, **model_kwargs):
520- super().__init__(model_info, model_kwargs=model_kwargs)
518+ class ForecastPipeline(BasicPipeline):
519+ def __init__(self, model_info, **model_kwargs):
520+ super().__init__(model_info, model_kwargs=model_kwargs)
521521
522- def preprocess(self, inputs: list[dict[str, dict[str, torch.Tensor] | torch.Tensor]], **infer_kwargs):
523- """
524- 在将输入数据传递给模型进行推理之前进行预处理,验证输入数据的形状和类型。
522+ def preprocess(self, inputs: list[dict[str, dict[str, torch.Tensor] | torch.Tensor]], **infer_kwargs):
523+ """
524+ 在将输入数据传递给模型进行推理之前进行预处理,验证输入数据的形状和类型。
525525
526- Args:
527- inputs (list[dict]):
528- 输入数据,字典列表,每个字典包含:
529- - 'targets': 形状为 (input_length,) 或 (target_count, input_length) 的张量。
530- - 'past_covariates': 可选,张量字典,每个张量形状为 (input_length,)。
531- - 'future_covariates': 可选,张量字典,每个张量形状为 (input_length,)。
526+ Args:
527+ inputs (list[dict]):
528+ 输入数据,字典列表,每个字典包含:
529+ - 'targets': 形状为 (input_length,) 或 (target_count, input_length) 的张量。
530+ - 'past_covariates': 可选,张量字典,每个张量形状为 (input_length,)。
531+ - 'future_covariates': 可选,张量字典,每个张量形状为 (input_length,)。
532532
533- infer_kwargs (dict, optional): 推理的额外关键字参数,如:
534- - ` output_length` (int): 如果提供'future_covariates',用于验证其有效性。
533+ infer_kwargs (dict, optional): 推理的额外关键字参数,如:
534+ - ` output_length` (int): 如果提供'future_covariates',用于验证其有效性。
535535
536- Raises:
537- ValueError: 如果输入格式不正确(例如,缺少键、张量形状无效)。
536+ Raises:
537+ ValueError: 如果输入格式不正确(例如,缺少键、张量形状无效)。
538538
539- Returns:
540- 经过预处理和验证的输入数据,可直接用于模型推理。
541- """
542- pass
539+ Returns:
540+ 经过预处理和验证的输入数据,可直接用于模型推理。
541+ """
542+ pass
543543
544- def forecast(self, inputs, **infer_kwargs):
545- """
546- 对给定输入执行预测。
544+ def forecast(self, inputs, **infer_kwargs):
545+ """
546+ 对给定输入执行预测。
547547
548- Parameters:
549- inputs: 用于进行预测的输入数据。类型和结构取决于模型的具体实现。
550- **infer_kwargs: 额外的推理参数,例如:
551- - ` output_length` (int): 模型应该生成的时间点数量。
548+ Parameters:
549+ inputs: 用于进行预测的输入数据。类型和结构取决于模型的具体实现。
550+ **infer_kwargs: 额外的推理参数,例如:
551+ - ` output_length` (int): 模型应该生成的时间点数量。
552552
553- Returns:
554- 预测输出,具体形式取决于模型的具体实现。
555- """
556- pass
553+ Returns:
554+ 预测输出,具体形式取决于模型的具体实现。
555+ """
556+ pass
557557
558- def postprocess(self, outputs: list[torch.Tensor], **infer_kwargs) -> list[torch.Tensor]:
559- """
560- 在推理后对模型输出进行后处理,验证输出数据的形状并确保其符合预期维度。
558+ def postprocess(self, outputs: list[torch.Tensor], **infer_kwargs) -> list[torch.Tensor]:
559+ """
560+ 在推理后对模型输出进行后处理,验证输出数据的形状并确保其符合预期维度。
561561
562- Args:
563- outputs:
564- 模型输出,2D张量列表,每个张量形状为 ` [target_count, output_length]` 。
562+ Args:
563+ outputs:
564+ 模型输出,2D张量列表,每个张量形状为 ` [target_count, output_length]` 。
565565
566- Raises:
567- InferenceModelInternalException: 如果输出张量形状无效(例如,维数错误)。
568- ValueError: 如果输出格式不正确。
566+ Raises:
567+ InferenceModelInternalException: 如果输出张量形状无效(例如,维数错误)。
568+ ValueError: 如果输出格式不正确。
569569
570- Returns:
571- list[torch.Tensor]:
572- 后处理后的输出,将是一个2D张量列表。
573- """
574- pass
575- ` ` `
570+ Returns:
571+ list[torch.Tensor]:
572+ 后处理后的输出,将是一个2D张量列表。
573+ """
574+ pass
575+ ` ` `
5765763 . 修改模型配置文件 config .json ,确保包含以下字段:
577577 ` ` ` JSON
578578 {
@@ -585,13 +585,13 @@ IoTDB> show models
585585 }
586586 ` ` `
587587
588- 1 . 必须通过 auto\_map 指定模型的 Config 类和模型类;
589- 2 . 必须集成并指定推理流水线类;
590- 3 . 对于 AINode 管理的内置(builtin)和自定义(user\_defined)模型,模型类别(model\_type)也作为不可重复的唯一标识。即,要注册的模型类别不得与任何已存在的模型类型重复,通过微调创建的模型将继承原模型的模型类别。
588+ * 必须通过 auto\_map 指定模型的 Config 类和模型类;
589+ * 必须集成并指定推理流水线类;
590+ * 对于 AINode 管理的内置(builtin)和自定义(user\_defined)模型,模型类别(model\_type)也作为不可重复的唯一标识。即,要注册的模型类别不得与任何已存在的模型类型重复,通过微调创建的模型将继承原模型的模型类别。
5915914 . 确保要注册的模型目录包含以下文件,且模型配置文件名称和权重文件名称不支持自定义:
592- 1 . 模型配置文件:config .json ;
593- 2 . 模型权重文件:model .safetensors ;
594- 3 . 模型代码:其它 .py 文件。
592+ * 模型配置文件:config .json ;
593+ * 模型权重文件:model .safetensors ;
594+ * 模型代码:其它 .py 文件。
595595
596596** 注册自定义模型的 SQL 语法如下所示:**
597597
0 commit comments