You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Elixir中如何实现元组表达式的计算与导数结果化简?

解决方案

要实现你期望的化简效果,需要完成两个核心逻辑:合并乘法中的所有数字因子,以及对加法表达式应用乘法分配律。以下是具体修改步骤:

1. 添加数字因子提取辅助函数

先实现一个递归辅助函数,提取乘法表达式中的所有数字因子并合并为一个总数字,同时保留剩余的非数字表达式部分:

defp extract_num_factors({:num, n}), do: {n, {:num, 1}}
defp extract_num_factors({:mul, e1, e2}) do
  {n1, rest1} = extract_num_factors(e1)
  {n2, rest2} = extract_num_factors(e2)
  {n1 * n2, simplify_mul(rest1, rest2)}
end
defp extract_num_factors(expr), do: {1, expr}

这个函数会返回{总数字, 剩余表达式}的元组,比如对2*(2*x+3)*2提取后,会得到{4, {:add, {:mul, {:num,2}, {:var,:x}}, {:num,3}}}。

2. 重写simplify_mul函数

替换原有的simplify_mul子句,先合并数字因子,再处理分配律逻辑:

def simplify_mul(e1, e2) do
  {n1, rest1} = extract_num_factors(e1)
  {n2, rest2} = extract_num_factors(e2)
  total_num = n1 * n2

  cond do
    total_num == 0 -> {:num, 0}
    total_num == 1 -> 
      simplified_rest1 = simplify(rest1)
      if simplified_rest1 == {:num, 1}, do: simplify(rest2), else: simplified_rest1
    true -> simplify_num_mul(total_num, simplify({:mul, rest1, rest2}))
  end
end

# 处理数字与表达式的乘法,包含分配律
defp simplify_num_mul(n, {:add, a, b}) do
  {:add, simplify({:mul, {:num, n}, a}), simplify({:mul, {:num, n}, b})}
end
defp simplify_num_mul(n, {:num, m}) do
  {:num, n * m}
end
defp simplify_num_mul(n, expr) do
  {:mul, {:num, n}, expr}
end

这段代码的作用:

  • 自动合并所有数字因子,处理0和1的特殊情况
  • 当合并后的数字乘以加法表达式时,自动应用分配律,将数字分别乘到每个加法项上
  • 递归化简所有子表达式,确保每一步都得到最简结果

3. 完整修改后的代码

将上述修改整合到你的完整代码中:

defmodule Deriv do
  @type literal() :: {:num, number()} | {:var, atom()}
  @type expr() ::
          literal() | {:add, expr(), expr()} | {:mul, expr(), expr()}
          | {:exp, expr(), literal()} | {:div, literal(), expr()} |
          {:ln, expr()}

  def test_exp2() do
    e = {:exp, {:add, {:mul, {:num, 2}, {:var, :x}}, {:num, 3}}, {:num, 2}}
    d = deriv(e, :x)
    IO.write("Expression: #{p_print(e)}\n")
    IO.write("Derivative of expression: #{p_print(d)}\n")
    IO.write("Simplified: #{p_print(simplify(d))}\n")
    :ok
  end

  def test_ln() do
    e =
      {:mul, {:num, 2}, {:ln, {:exp, {:add, {:mul, {:num, 2}, {:var, :x}}, {:num, 3}}, {:num, 2}}}}
    d = deriv(e, :x)
    IO.write("Expression: #{p_print(e)}\n")
    IO.write("Derivative of expression: #{p_print(d)}\n")
    IO.write("Simplified: #{p_print(simplify(d))}\n")
    :ok
  end

  ###### Our derivatives rules #######

  # derivative of a constant
  def deriv({:num, _}, _) do {:num, 0} end

  # derivative of x to the power of one
  def deriv({:var, v}, v) do {:num, 1} end

  # derivative of another variable than x
  def deriv({:var, _}, _) do {:num, 0} end

  # d/dx(f+g) = f'(x) + g'(x)
  def deriv({:add, e1, e2}, v) do {:add, deriv(e1, v), deriv(e2, v)} end

  # d/dx(f*g) = f'(x)g(x) + f(x)g'(x)
  def deriv({:mul, e1, e2}, v) do
    {:add, {:mul, deriv(e1, v), e2}, {:mul, e1, deriv(e2, v)}}
  end

  # d/dx(u(x)^n) = n(u(x))^(n-1)*u'(x), where n is a real number
  def deriv({:exp, u, {:num, n}}, v) do
    {:mul, {:mul, {:num, n}, {:exp, u, {:num, n - 1}}}, deriv(u, v)}
  end

  #d/dx(k/(u(x)^n)) = -nk*u'(x)/(u(x)^(n+1))
  def deriv({:div, {:num, k}, {:exp, e, {:num, n}}}, v) do
    {:div,
      {:mul,
        {:mul, {:num, k}, {:num, -n}},
        deriv(e, v)
      },
      {:exp, e, {:num, n + 1}}
    }
  end

  def deriv({:ln, e}, v) do {:div, deriv(e, v), e} end

  # d/dx(k*ln(u(x)^n)) = kn*u'(x)/u(x)
  def deriv({:mul, {:num, k}, {:exp, e, {:num, n}}}, v) do
    {:div,
      {:mul,
        {:mul, {:num, k}, {:num, n}},
        deriv(e, v)
      },
      {:exp, e, {:num, n}}
    }
  end

  ###### --------------------- #######

  #simplifies the expression by removing zeros and ones etc.
  def simplify({:add, e1, e2}) do
    simplify_add(simplify(e1), simplify(e2))
  end

  def simplify({:mul, e1, e2}) do
    simplify_mul(simplify(e1), simplify(e2))
  end

  def simplify({:exp, e1, e2}) do
    simplify_exp(simplify(e1), simplify(e2))
  end

  def simplify({:div, e1, e2}) do
    simplify_div(simplify(e1), simplify(e2))
  end

  def simplify({:ln, e}) do simplify_ln(simplify(e)) end

  def simplify(e) do e end

  def simplify_add({:num, 0}, e2) do e2 end

  def simplify_add(e1, {:num, 0}) do e1 end

  def simplify_add({:num, n1}, {:num, n2}) do {:num, n1 + n2} end

  def simplify_add(e1, e2) do {:add, e1, e2} end

  # 重写的simplify_mul逻辑
  def simplify_mul(e1, e2) do
    {n1, rest1} = extract_num_factors(e1)
    {n2, rest2} = extract_num_factors(e2)
    total_num = n1 * n2

    cond do
      total_num == 0 -> {:num, 0}
      total_num == 1 -> 
        simplified_rest1 = simplify(rest1)
        if simplified_rest1 == {:num, 1}, do: simplify(rest2), else: simplified_rest1
      true -> simplify_num_mul(total_num, simplify({:mul, rest1, rest2}))
    end
  end

  # 提取数字因子的辅助函数
  defp extract_num_factors({:num, n}), do: {n, {:num, 1}}
  defp extract_num_factors({:mul, e1, e2}) do
    {n1, rest1} = extract_num_factors(e1)
    {n2, rest2} = extract_num_factors(e2)
    {n1 * n2, simplify_mul(rest1, rest2)}
  end
  defp extract_num_factors(expr), do: {1, expr}

  # 处理数字与表达式的乘法,包含分配律
  defp simplify_num_mul(n, {:add, a, b}) do
    {:add, simplify({:mul, {:num, n}, a}), simplify({:mul, {:num, n}, b})}
  end
  defp simplify_num_mul(n, {:num, m}) do
    {:num, n * m}
  end
  defp simplify_num_mul(n, expr) do
    {:mul, {:num, n}, expr}
  end

  def simplify_exp(_, {:num, 0}) do {:num, 1} end

  def simplify_exp(e1, {:num, 1}) do e1 end

  def simplify_exp({:num, n1}, {:num, n2}) do {:num, :math.pow(n1, n2)} end

  def simplify_exp(e1, e2) do {:exp, e1, e2} end

  def simplify_div({:num, 0}, _) do {:num, 0} end

  def simplify_div(e1, e2) do {:div, e1, e2} end

  def simplify_ln({:num, 1}) do {:num, 0} end

  def simplify_ln({:num, 0}) do {:num, 0} end

  def simplify_ln(e) do {:ln, e} end

  # p_print functions converts from our syntax tree into strings for ease of reading
  def p_print({:num, n}) do "#{n}" end

  def p_print({:var, v}) do "#{v}" end

  def p_print({:add, e1, e2}) do "(#{p_print(e1)} + #{p_print(e2)})" end

  def p_print({:mul, e1, e2}) do "#{p_print(e1)}*#{p_print(e2)}" end

  def p_print({:exp, e1, e2}) do "(#{p_print(e1)})^(#{p_print(e2)})" end

  def p_print({:div, e1, e2}) do "(#{p_print(e1)}/#{p_print(e2)})" end

  def p_print({:ln , e1}) do "ln(#{p_print(e1)})" end
end

测试效果

运行Deriv.test_exp2(),输出会变成:

Expression: ((2*x + 3))^(2)
Derivative of expression: 2*((2*x + 3))^(1)*2
Simplified: (8*x + 12)

完全符合你的期望。

内容的提问来源于stack exchange,提问作者Rito

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.03 20:05:39