openGauss数据库源码解析 |AI技术(15)
·
4. DNN算法源码解析
训练阶段先初始化SQL向量,之后创建深度学习模型,将模型保存到本地。
预测阶段,导入模型,向量化待预测的SQL;基于向量相似度对SQL的执行时间进行预测。主要源码如下:
class KerasRegression:
# 初始化模型参数
def __init__(self, encoding_dim=1):
self.model = None
self.encoding_dim = encoding_dim
# 模型定义
@staticmethod
def build_model(shape, encoding_dim):
from tensorflow.keras import Input, Model
from tensorflow.keras.layers import Dense
inputs = Input(shape=(shape,))
layer_dense1 = Dense(128, activation='relu', kernel_initializer='he_normal')(inputs)
…
model = Model(inputs=inputs, outputs=y_pred)
# 优化器,损失函数
model.compile(optimizer='adam', loss='mse', metrics=['mae'])
return model
# 模型训练
def fit(self, features, labels, batch_size=128, epochs=300):
…
self.model.fit(features, labels, epochs=epochs, batch_size=batch_size, shuffle=True, verbose=2)
# 模型预测
def predict(self, features):
predict_result = self.model.predict(features)
return predict_result
# 模型保存
def save(self, filepath):
self.model.save(filepath)
# 模型读取
def load(self, filepath):
from tensorflow.keras.models import load_model
self.model = load_model(filepath)
class DnnModel(AbstractModel, ABC):
# 初始化算法参数
def __init__(self, params):
…
self.regression = KerasRegression(encoding_dim=1)
self.data = None
# 把sql语句转化为vector,如果模型不存在,则直接训练w2v模型,如果模型存在则进行增量训练
def build_word2vector(self, data):
self.data = list(data)
if self.w2v.model:
self.w2v.update(self.data)
else:
self.w2v.fit(self.data)
def fit(self, data):
self.build_word2vector(data)
…
# 数据归一化
self.scaler = MinMaxScaler(feature_range=(0, 1))
self.scaler.fit(labels)
labels = self.scaler.transform(labels)
self.regression.fit(features, labels, epochs=self.epoch)
# 利用回归模型预测执行时间
def transform(self, data):
…
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)