【发布时间】:2014-12-17 21:32:09
【问题描述】:
我正在试用linear regression with automatic differentiation 的附加代码。 它指定由两个浮点数组成的数据类型 [Dual][2],并声明它是 Num、Fractional 和 Floating 的实例。 与所有拟合/回归任务一样,有一个标量成本函数,由拟合参数 c 和 m 进行参数化,还有一个优化器通过梯度下降改进这两个参数的估计。
问题 我正在使用 GHC 7.8.3,作者明确提到这是 H98 代码(我在标题中提到它是因为这是我能想到的我的设置和作者的设置之间的唯一实质性区别,但是如果错误请纠正)。 为什么它会在成本函数的定义中窒息? 我的理解是:函数 idD 和 constD 将 Floats 映射到 Duals,g 是多态的(它可以对 Dual 输入执行代数运算,因为 Dual 继承自 Num、Fractional 和 Floating),并且 deriv 将 Duals 映射到 Doubles。 推断出 g 的类型签名(数据的 eta 减少成本函数)。我尝试省略它,并通过将 Floating 类型约束替换为 Fractional 来使其更通用。 此外,我尝试将 c 和 m 的数字类型与 (fromIntegral c :: Double) 内联,但无济于事。
特别是这段代码给出了这个错误:
No instance for (Integral Dual) arising from a use of ‘g’
In the first argument of ‘flip’, namely ‘g’
In the expression: flip g (constD c)
In the second argument of ‘($)’, namely ‘flip g (constD c) $ idD m’
有什么提示吗?我敢肯定这是一个非常菜鸟的问题,但我就是不明白。
完整代码如下:
{-# LANGUAGE NoMonomorphismRestriction #-}
module ADfw (Dual(..), f, idD, cost) where
data Dual = Dual Double Double deriving (Eq, Show)
constD :: Double -> Dual
constD x = Dual x 0
idD :: Double -> Dual
idD x = Dual x 1.0
instance Num Dual where
fromInteger n = constD $ fromInteger n
(Dual x x') + (Dual y y') = Dual (x+y) (x' + y')
(Dual x x') * (Dual y y') = Dual (x*y) (x*y' + y*x')
negate (Dual x x') = Dual (negate x) (negate x')
signum _ = undefined
abs _ = undefined
instance Fractional Dual where
fromRational p = constD $ fromRational p
recip (Dual x x') = Dual (1.0 / x) (- x' / (x*x))
instance Floating Dual where
pi = constD pi
exp (Dual x x') = Dual (exp x) (x' * exp x)
log (Dual x x') = Dual (log x) (x' / x)
sqrt (Dual x x') = Dual (sqrt x) (x' / (2 * sqrt x))
sin (Dual x x') = Dual (sin x) (x' * cos x)
cos (Dual x x') = Dual (cos x) (x' * (- sin x))
sinh (Dual x x') = Dual (sinh x) (x' * cosh x)
cosh (Dual x x') = Dual (cosh x) (x' * sinh x)
asin (Dual x x') = Dual (asin x) (x' / sqrt (1 - x*x))
acos (Dual x x') = Dual (acos x) (x' / (-sqrt (1 - x*x)))
atan (Dual x x') = Dual (atan x) (x' / (1 + x*x))
asinh (Dual x x') = Dual (asinh x) (x' / sqrt (1 + x*x))
acosh (Dual x x') = Dual (acosh x) (x' / (sqrt (x*x - 1)))
atanh (Dual x x') = Dual (atanh x) (x' / (1 - x*x))
-- example
-- f = sqrt . (* 3) . sin
-- f' x = 3 * cos x / (2 * sqrt (3 * sin x))
-- linear fit sum-of-squares cost
-- cost :: Fractional s => s -> s -> [s] -> [s] -> s
cost m c x y = (/ (2 * (fromIntegral $ length x))) $
sum $ zipWith errSq x y
where
errSq xi yi = zi * zi
where
zi = yi - (m * xi + c)
-- test data
x_ = [1..10]
y_ = [a | a <- [1..20], a `mod` 2 /= 0]
-- learning rate
gamma = 0.04
g :: (Integral s, Fractional s) => s -> s -> s
g m c = cost m c x_ y_
deriv (Dual _ x') = x'
z_ = (0.1, 0.1) : map h z_
h (c, m) = (c - gamma * cd, m - gamma * md) where
cd = deriv $ g (constD m) $ idD c
md = deriv $ flip g (constD c) $ idD m
-- check for convergence
main = do
take 2 $ drop 1000 $ map (\(c, m) -> cost m c x_ y_) z_
take 2 $ drop 1000 $ z_
其中测试数据 x_ 和 y_ 是数组,学习率 gamma 是一个标量。
[2]:如果我们将导数视为运算符,则对偶对象的两个字段实际上是彼此相邻的
【问题讨论】:
-
如果有人偶然发现这个例子,错误在于使用列表推导生成数据 x_ 和 y_。如果您明确声明他们的条目,例如x_=[1,2,4,10,2] 等,它只是工作
标签: haskell