Tensorflow 2.0 常见函数用法(一)

文章目录

  • [0. 基础用法](#0. 基础用法)
  • [1. tf.cast](#1. tf.cast)
  • [2. tf.keras.layers.Dense](#2. tf.keras.layers.Dense)
  • [3. tf.variable_scope](#3. tf.variable_scope)
  • [4. tf.squeeze](#4. tf.squeeze)
  • [5. tf.math.multiply](#5. tf.math.multiply)

0. 基础用法

Tensorflow 的用法不定期更新遇到的一些用法,之前已经包含了基础用法参考这里 ,具体包含如下图的方法:

本文介绍其他常见的方法。

1. tf.cast

张量类型强制转换

官方用法:

python 复制代码
tf.cast(
    x, dtype, name=None
)

示例:

python 复制代码
x = tf.constant([1.8, 2.2], dtype=tf.float32)
print(tf.cast(x, tf.int32))

# 输出
tf.Tensor([1 2], shape=(2,), dtype=int32)

2. tf.keras.layers.Dense

构建一个全连接层

在1.0中是 tf.layers.dense ,2.0中可以用下面方法兼容:

python 复制代码
import tensorflow.compat.v1 as tf
tf.layers.dense(xxx)

官方用法:

python 复制代码
tf.keras.layers.Dense(
    units,
    activation=None,
    use_bias=True,
    kernel_initializer='glorot_uniform',
    bias_initializer='zeros',
    kernel_regularizer=None,
    bias_regularizer=None,
    activity_regularizer=None,
    kernel_constraint=None,
    bias_constraint=None,
    **kwargs
)

3. tf.variable_scope

这是 v1 版本的用法,用于管理变量

官方用法:

python 复制代码
tf.compat.v1.variable_scope(
    name_or_scope,
    default_name=None,
    values=None,
    initializer=None,
    regularizer=None,
    caching_device=None,
    partitioner=None,
    custom_getter=None,
    reuse=None,
    dtype=None,
    use_resource=None,
    constraint=None,
    auxiliary_name_scope=True
)

示例:

python 复制代码
import tensorflow as tf
with tf.variable_scope("one"):
    o=tf.get_variable("f",[1])
with tf.variable_scope("two"):
    o1=tf.get_variable("f",[1])

# 抛错,因为变量的作用范围不一样
# 一个作用域是one/f,一个作用域是two/f
assert o == o1

4. tf.squeeze

从张量的形状中移除大小为1的维度。该函数返回一个张量,这个张量是将原始input中所有维度为1的那些维都删掉的结果。

axis 可以用来指定要删掉的为1的维度,此处要注意指定的维度必须确保其是1,否则会报错。

官方用法:

python 复制代码
tf.squeeze(
    input, axis=None, name=None
)

示例:

python 复制代码
# 注意,a的shape是1*6,即存在一个大小为1的维度
a = tf.constant([1, 2, 3, 4, 5, 6], shape=[1, 6])
print(a)
b = tf.squeeze(a, [0]) # 删除第0个维度为1的
# b = tf.squeeze(a) 的结果是一样的
print(b)

# 输出
tf.Tensor([[1 2 3 4 5 6]], shape=(1, 6), dtype=int32)
tf.Tensor([1 2 3 4 5 6], shape=(6,), dtype=int32)
python 复制代码
a = tf.constant([1, 2, 3, 4, 5, 6], shape=[6, 1])
print(a)
b = tf.squeeze(a, [1])
print(b)

# 输出
tf.Tensor(
[[1]
 [2]
 [3]
 [4]
 [5]
 [6]], shape=(6, 1), dtype=int32)
tf.Tensor([1 2 3 4 5 6], shape=(6,), dtype=int32)
python 复制代码
a = tf.constant([1, 2, 3, 4, 5, 6], shape=[2, 3])
print(a)
b = tf.squeeze(a) # 如果不存在大小为1的维度,那么保持不变
print(b)

# 输出
tf.Tensor(
[[1 2 3]
 [4 5 6]], shape=(2, 3), dtype=int32)
tf.Tensor(
[[1 2 3]
 [4 5 6]], shape=(2, 3), dtype=int32)

5. tf.math.multiply

元素相乘

在 1.0 中是 tf.multiply
官方用法:

python 复制代码
tf.math.multiply(
    x, y, name=None
)

示例:

python 复制代码
a = tf.constant([1, 2, 3, 4, 5, 6], shape=[2, 3])
print(tf.multiply(a, 2))
print(tf.multiply(a, a))

# 输出
tf.Tensor(
[[ 2  4  6]
 [ 8 10 12]], shape=(2, 3), dtype=int32)

tf.Tensor(
[[ 1  4  9]
 [16 25 36]], shape=(2, 3), dtype=int32)
python 复制代码
x = tf.ones([1, 2]);
y = tf.ones([2, 1]);
print(x * y)  # Taking advantage of operator overriding
print(tf.multiply(x, y))

# 输出,如果维度不一致,会尝试匹配维度
tf.Tensor(
[[1. 1.]
 [1. 1.]], shape=(2, 2), dtype=float32)
tf.Tensor(
[[1. 1.]
 [1. 1.]], shape=(2, 2), dtype=float32)
相关推荐
火山引擎开发者社区6 小时前
技术速递|使用 GitHub Copilot CLI 构建 Emoji 列表生成器
人工智能
weelinking6 小时前
【产品】12_接入数据库——让数据永久保存
jvm·数据库·python·react.js·数据挖掘·前端框架·产品经理
codefan※6 小时前
干掉“幻觉“实战:如何构建企业级知识图谱增强 RAG
人工智能·知识图谱
wukangjupingbb6 小时前
传统基于药物 SMILES 序列和蛋白质氨基酸序列的 DTI(Drug-Target Interaction)预测方法的缺陷
人工智能
沪漂阿龙7 小时前
Codex 额度重置周期变化:AI 编程免费试玩时代正在结束
人工智能
程序大视界7 小时前
【Python系列课程】Python正则表达式(下):环视、命名分组与日志实战
开发语言·python·正则表达式
TickDB7 小时前
美股行情 API 接入避坑:REST 快照、WebSocket 推送、盘前盘后数据的边界
人工智能·python·websocket·行情数据 api
装不满的克莱因瓶7 小时前
深入理解卷积神经网络(CNN)——从原理到代码实践
人工智能·神经网络·cnn
完成大叔7 小时前
模块二,Agent知识图谱的工具链思考
人工智能
lauo7 小时前
ibbot手机发布:搭载poplang技术 + token节点经济,革新AI手机体验
人工智能·智能手机